diff --git a/.github/workflows/notify-on-merge.yml b/.github/workflows/notify-on-merge.yml new file mode 100644 index 000000000..99be4edf0 --- /dev/null +++ b/.github/workflows/notify-on-merge.yml @@ -0,0 +1,30 @@ +name: Notify on Merge + +on: + push: + branches: + - master + +jobs: + notify: + runs-on: ubuntu-latest + if: github.repository == 'Comfy-Org/ComfyUI' + steps: + - name: Notify downstream + env: + DISPATCH_TOKEN: ${{ secrets.SYNC_DISPATCH_TOKEN }} + TARGET_REPO: ${{ secrets.SYNC_TARGET_REPO }} + COMMIT_SHA: ${{ github.sha }} + run: | + set -euo pipefail + if [ -z "${DISPATCH_TOKEN:-}" ] || [ -z "${TARGET_REPO:-}" ]; then + echo "::notice::SYNC_DISPATCH_TOKEN/SYNC_TARGET_REPO not set; skipping downstream notify." + exit 0 + fi + PAYLOAD="$(jq -n --arg sha "$COMMIT_SHA" \ + '{ event_type: "upstream-push", client_payload: { sha: $sha } }')" + curl -fsSL --connect-timeout 10 --max-time 60 -X POST \ + -H "Accept: application/vnd.github+json" \ + -H "Authorization: Bearer ${DISPATCH_TOKEN}" \ + "https://api.github.com/repos/${TARGET_REPO}/dispatches" \ + -d "$PAYLOAD" diff --git a/README.md b/README.md index 8830b62ce..118206d1e 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,7 @@ ComfyUI is the AI creation engine for visual professionals who demand control over every model, every parameter, and every output. Its powerful and modular node graph interface empowers creatives to generate images, videos, 3D models, audio, and more... - ComfyUI natively supports the latest open-source state of the art models. -- API nodes provide access to the best closed source models such as Nano Banana, Seedance, Hunyuan3D, etc. +- [Partner nodes](https://docs.comfy.org/tutorials/partner-nodes/overview#partner-nodes) provide access to the best closed source models such as Nano Banana, Seedance, Hunyuan3D, etc. - It is available on Windows, Linux, and macOS, locally with our [desktop application](https://www.comfy.org/download), our [portable install](#installing) or on our [cloud](https://www.comfy.org/cloud). - The most sophisticated workflows can be exposed through a simple UI thanks to App Mode. - It integrates seamlessly into production pipelines with our API endpoints. @@ -74,7 +74,7 @@ See what ComfyUI can do with the [newer template workflows](https://comfy.org/wo - [Image editing](https://comfy.org/workflows/tag/image-edit/): Flux Kontext, Flux.2 Klein, Qwen Image Edit, HiDream E1.1 and O1, OmniGen2, Boogu, JoyImage Edit, MageFlow Edit, and LongCat Image Edit. - [Video generation](https://comfy.org/workflows/tag/video-generation/): Wan 2.1 and 2.2, LTX-Video 2 and 2.3, HunyuanVideo 1.5, Kandinsky 5 Video, CogVideoX, Cosmos Predict2, Bernini-R, SCAIL 2, and Mochi. - [Audio and video generation](https://comfy.org/workflows/): MiniMax H3 and LTX-AV. - - [Audio generation](https://comfy.org/workflows/tag/text-to-audio/): ACE-Step 1.5 and Stable Audio 3. + - [Audio generation](https://comfy.org/workflows/tag/text-to-audio/): ACE-Step 1.5, Stable Audio 3 and MiniMax Music 3 - [3D and vision](https://comfy.org/workflows/): Hunyuan3D 2.1, TripoSplat, SeedVR2, SUPIR, Depth Anything 3, MoGe, SAM 3 and 3.1, RT-DETRv4, and BiRefNet. - [Text generation](https://comfy.org/workflows/tag/text-generation/): Gemma 3 and 4, Qwen3, Qwen3.5, and Qwen3-VL, including multimodal inputs. - Load complete checkpoints or separate diffusion models, VAEs, text encoders, LoRAs, ControlNets, adapters, and upscalers from supported model formats. @@ -194,7 +194,7 @@ Python 3.14 works but some custom nodes may have issues. The free threaded varia Python 3.13 is very well supported. If you have trouble with some custom node dependencies on 3.13 you can try 3.12 -torch 2.5 is minimally supported but using a newer version is extremely recommended. Some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. If your pytorch is more than 6 months old, please update it. +torch 2.7 is minimally supported but using a newer version is extremely recommended. Using a cu130 or above version of pytorch is required on Nvidia 20 series and above. Some features and optimizations might only work on newer versions. We generally recommend using the latest major version of pytorch with the latest cuda version unless it is less than 2 weeks old. If your pytorch is more than 6 months old, please update it. ### Instructions: diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index e25b8a57f..62fb97025 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -2,6 +2,7 @@ import asyncio import functools import json import logging +import mimetypes import os import urllib.parse import uuid @@ -18,7 +19,7 @@ from app.assets.api.schemas_in import ( AssetValidationError, UploadError, ) -from app.assets.helpers import validate_blake3_hash +from app.assets.helpers import normalize_tags, validate_blake3_hash from app.assets.api.upload import ( delete_temp_file_if_exists, parse_multipart_upload, @@ -32,6 +33,7 @@ from app.assets.services import ( create_from_hash, delete_asset_reference, get_asset_detail, + get_preview_file_paths, list_assets_page, list_tags, remove_tags, @@ -40,7 +42,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.path_utils import compute_asset_response_paths from app.assets.services.tagging import list_tag_histogram ROUTES = web.RouteTableDef() @@ -117,6 +119,87 @@ def _build_validation_error_response(code: str, ve: ValidationError) -> web.Resp return _build_error_response(400, code, "Validation failed.", {"errors": errors}) +class InvalidTagFilterError(Exception): + """Invalid combination of tag-filter query parameters.""" + + def __init__(self, message: str, details: dict): + super().__init__(message) + self.details = details + + +# Caps the per-tag EXISTS fan-out; deliberately covers the legacy spellings too. +MAX_TAG_FILTER_TAGS = 100 + + +def _resolve_tag_filters( + q: schemas_in.ListAssetsQuery | schemas_in.TagsRefineQuery, +) -> tuple[list[str], list[str], list[str]]: + """Resolve legacy (include/exclude) and new (all/any/none) tag-filter + spellings into effective (all, any, none) lists. + + Combination validation applies only when the request uses at least one + new-name parameter (non-empty after normalisation); requests using only + the legacy names keep their historical behaviour, including degenerate + combinations like include_tags=a&exclude_tags=a. + """ + # model_dump, not attribute access: deprecated fields warn on every attribute read. + legacy = q.model_dump(include={"include_tags", "exclude_tags"}) + include_tags = normalize_tags(legacy["include_tags"]) + exclude_tags = normalize_tags(legacy["exclude_tags"]) + tags_all = normalize_tags(q.tags_all) + tags_any = normalize_tags(q.tags_any) + tags_none = normalize_tags(q.tags_none) + + for param_name, values in ( + ("include_tags", include_tags), + ("exclude_tags", exclude_tags), + ("tags_all", tags_all), + ("tags_any", tags_any), + ("tags_none", tags_none), + ): + if len(values) > MAX_TAG_FILTER_TAGS: + raise InvalidTagFilterError( + f"'{param_name}' lists {len(values)} tags; the maximum is " + f"{MAX_TAG_FILTER_TAGS}.", + { + "parameter": param_name, + "count": len(values), + "max": MAX_TAG_FILTER_TAGS, + }, + ) + + if not (tags_all or tags_any or tags_none): + return include_tags, [], exclude_tags + + if include_tags and tags_all: + raise InvalidTagFilterError( + "Cannot combine 'include_tags' with 'tags_all'; use 'tags_all'.", + {"parameters": ["include_tags", "tags_all"]}, + ) + if exclude_tags and tags_none: + raise InvalidTagFilterError( + "Cannot combine 'exclude_tags' with 'tags_none'; use 'tags_none'.", + {"parameters": ["exclude_tags", "tags_none"]}, + ) + + all_param, all_list = ( + ("tags_all", tags_all) if tags_all else ("include_tags", include_tags) + ) + none_param, none_list = ( + ("tags_none", tags_none) if tags_none else ("exclude_tags", exclude_tags) + ) + + conflicting = sorted(set(all_list) & set(none_list)) + if conflicting: + raise InvalidTagFilterError( + f"Query can never match: {', '.join(repr(t) for t in conflicting)} " + f"required by '{all_param}' but rejected by '{none_param}'.", + {"conflicting_tags": conflicting, "parameters": [all_param, none_param]}, + ) + + return all_list, tags_any, none_list + + def _validate_sort_field(requested: str | None) -> str: if not requested: return "created_at" @@ -126,44 +209,62 @@ def _validate_sort_field(requested: str | None) -> str: return "created_at" -def _build_preview_url_from_view(tags: list[str], user_metadata: dict[str, Any] | None) -> str | None: - """Build a /api/view preview URL from asset tags and user_metadata filename.""" - if not user_metadata: +# What a client can render from the bytes themselves; anything else needs a nominated preview. +PREVIEWABLE_MIME_PREFIXES = ("image/", "video/", "audio/", "text/") + +# models is deliberately absent: /api/view has no directory type for it. +VIEWABLE_NAMESPACES = frozenset({"input", "output", "temp"}) + + +def _has_previewable_content(asset: schemas.AssetData | None, file_path: str | None) -> bool: + if asset is None: + return False + # Resolved from the path, not the caller-editable name, so a rename cannot change what previews. + raw = asset.mime_type or mimetypes.guess_type(file_path or "")[0] or "" + return raw.split(";", 1)[0].strip().lower().startswith(PREVIEWABLE_MIME_PREFIXES) + + +def _build_view_url(file_path: str | None) -> str | None: + # /api/view is a FileResponse: byte-range seeking, no user header, no access write. + if not file_path: return None - filename = user_metadata.get("filename") - if not filename: + paths = compute_asset_response_paths(file_path) + if not paths: + return None + logical_path, relative_path = paths + namespace = logical_path.split("/", 1)[0] + if namespace not in VIEWABLE_NAMESPACES or not relative_path: return None - if "input" in tags: - view_type = "input" - elif "output" in tags: - view_type = "output" - else: - return None - - subfolder = "" - if "/" in filename: - subfolder, filename = filename.rsplit("/", 1) - - encoded_filename = urllib.parse.quote(filename, safe="") - url = f"/api/view?type={view_type}&filename={encoded_filename}" + subfolder, _, filename = relative_path.rpartition("/") + url = f"/api/view?type={namespace}&filename={urllib.parse.quote(filename, safe='')}" if subfolder: url += f"&subfolder={urllib.parse.quote(subfolder, safe='')}" return url -def _build_asset_response(result: schemas.AssetDetailResult | schemas.UploadResult) -> schemas_out.Asset: - """Build an Asset response from a service result.""" +def _resolve_preview_paths( + results: "list[schemas.AssetDetailResult] | list[schemas.AssetSummaryData]", +) -> dict[str, str]: + # A miss means no live preview - that is what keeps a soft-deleted one quiet. + preview_ids = {r.ref.preview_id for r in results if r.ref.preview_id} + return get_preview_file_paths(sorted(preview_ids)) + + +def _build_asset_response( + result: schemas.AssetDetailResult | schemas.UploadResult, + preview_paths: dict[str, str], +) -> schemas_out.Asset: if result.ref.preview_id: - preview_detail = get_asset_detail(result.ref.preview_id) - if preview_detail: - preview_url = _build_preview_url_from_view(preview_detail.tags, preview_detail.ref.user_metadata) - else: - preview_url = None + # A nominated preview is one whatever it holds, so no media check here. + preview_url = _build_view_url(preview_paths.get(result.ref.preview_id)) + elif _has_previewable_content(result.asset, result.ref.file_path): + preview_url = _build_view_url(result.ref.file_path) else: - preview_url = _build_preview_url_from_view(result.tags, result.ref.user_metadata) + preview_url = None if result.ref.file_path: - display_name = compute_display_name(result.ref.file_path) + paths = compute_asset_response_paths(result.ref.file_path) + display_name = paths[1] if paths else None # In-root loader path (model category dropped): what model loaders consume. loader_path = result.ref.loader_path else: @@ -217,6 +318,11 @@ async def list_assets_route(request: web.Request) -> web.Response: except ValidationError as ve: return _build_validation_error_response("INVALID_QUERY", ve) + try: + tags_all, tags_any, tags_none = _resolve_tag_filters(q) + except InvalidTagFilterError as e: + return _build_error_response(400, "INVALID_TAG_FILTER", str(e), e.details) + sort = _validate_sort_field(q.sort) order_candidate = (q.order or "desc").lower() order = order_candidate if order_candidate in {"asc", "desc"} else "desc" @@ -224,8 +330,9 @@ async def list_assets_route(request: web.Request) -> web.Response: try: result = list_assets_page( owner_id=USER_MANAGER.get_request_user_id(request), - include_tags=q.include_tags, - exclude_tags=q.exclude_tags, + include_tags=tags_all, + exclude_tags=tags_none, + any_tags=tags_any, name_contains=q.name_contains, metadata_filter=q.metadata_filter, limit=q.limit, @@ -237,7 +344,8 @@ async def list_assets_route(request: web.Request) -> web.Response: except InvalidCursorError as e: return _build_error_response(400, "INVALID_CURSOR", str(e)) - summaries = [_build_asset_response(item) for item in result.items] + preview_paths = _resolve_preview_paths(result.items) + summaries = [_build_asset_response(item, preview_paths) for item in result.items] # has_more semantics differ by mode: # - cursor mode: a non-empty next_cursor means there are more results. @@ -276,7 +384,7 @@ async def get_asset_route(request: web.Request) -> web.Response: {"id": reference_id}, ) - payload = _build_asset_response(result) + payload = _build_asset_response(result, _resolve_preview_paths([result])) except ValueError as e: return _build_error_response( 404, "ASSET_NOT_FOUND", str(e), {"id": reference_id} @@ -407,7 +515,7 @@ async def create_asset_from_hash_route(request: web.Request) -> web.Response: 404, "ASSET_NOT_FOUND", f"Asset content {body.hash} does not exist" ) - asset = _build_asset_response(result) + asset = _build_asset_response(result, _resolve_preview_paths([result])) payload_out = schemas_out.AssetCreated( **asset.model_dump(), created_new=result.created_new, @@ -498,7 +606,7 @@ async def upload_asset(request: web.Request) -> web.Response: logging.exception("upload_asset failed for owner_id=%s", owner_id) return _build_error_response(500, "INTERNAL", "Unexpected server error.") - asset = _build_asset_response(result) + asset = _build_asset_response(result, _resolve_preview_paths([result])) payload_out = schemas_out.AssetCreated( **asset.model_dump(), created_new=result.created_new, @@ -528,7 +636,7 @@ async def update_asset_route(request: web.Request) -> web.Response: owner_id=USER_MANAGER.get_request_user_id(request), preview_id=body.preview_id, ) - payload = _build_asset_response(result) + payload = _build_asset_response(result, _resolve_preview_paths([result])) except PermissionError as pe: return _build_error_response(403, "FORBIDDEN", str(pe), {"id": reference_id}) except ValueError as ve: @@ -715,10 +823,16 @@ async def get_tags_refine(request: web.Request) -> web.Response: except ValidationError as ve: return _build_validation_error_response("INVALID_QUERY", ve) + try: + tags_all, tags_any, tags_none = _resolve_tag_filters(q) + except InvalidTagFilterError as e: + return _build_error_response(400, "INVALID_TAG_FILTER", str(e), e.details) + tag_counts = list_tag_histogram( owner_id=USER_MANAGER.get_request_user_id(request), - include_tags=q.include_tags, - exclude_tags=q.exclude_tags, + include_tags=tags_all, + exclude_tags=tags_none, + any_tags=tags_any, name_contains=q.name_contains, metadata_filter=q.metadata_filter, limit=q.limit, diff --git a/app/assets/api/schemas_in.py b/app/assets/api/schemas_in.py index 38a942b7b..862700a24 100644 --- a/app/assets/api/schemas_in.py +++ b/app/assets/api/schemas_in.py @@ -50,8 +50,12 @@ class ParsedUpload: class ListAssetsQuery(BaseModel): - include_tags: list[str] = Field(default_factory=list) - exclude_tags: list[str] = Field(default_factory=list) + # Deprecated spellings: include_tags ≡ tags_all, exclude_tags ≡ tags_none. + include_tags: list[str] = Field(default_factory=list, deprecated=True) + exclude_tags: list[str] = Field(default_factory=list, deprecated=True) + tags_all: list[str] = Field(default_factory=list) + tags_any: list[str] = Field(default_factory=list) + tags_none: list[str] = Field(default_factory=list) name_contains: str | None = None # Accept either a JSON string (query param) or a dict @@ -70,7 +74,10 @@ class ListAssetsQuery(BaseModel): ) order: Literal["asc", "desc"] = "desc" - @field_validator("include_tags", "exclude_tags", mode="before") + @field_validator( + "include_tags", "exclude_tags", "tags_all", "tags_any", "tags_none", + mode="before", + ) @classmethod def _split_csv_tags(cls, v): # Accept "a,b,c" or ["a","b"] (we are liberal in what we accept) @@ -154,13 +161,20 @@ class CreateFromHashBody(BaseModel): class TagsRefineQuery(BaseModel): - include_tags: list[str] = Field(default_factory=list) - exclude_tags: list[str] = Field(default_factory=list) + # Deprecated spellings: include_tags ≡ tags_all, exclude_tags ≡ tags_none. + include_tags: list[str] = Field(default_factory=list, deprecated=True) + exclude_tags: list[str] = Field(default_factory=list, deprecated=True) + tags_all: list[str] = Field(default_factory=list) + tags_any: list[str] = Field(default_factory=list) + tags_none: list[str] = Field(default_factory=list) name_contains: str | None = None metadata_filter: dict[str, Any] | None = None limit: conint(ge=1, le=1000) = 100 - @field_validator("include_tags", "exclude_tags", mode="before") + @field_validator( + "include_tags", "exclude_tags", "tags_all", "tags_any", "tags_none", + mode="before", + ) @classmethod def _split_csv_tags(cls, v): if v is None: diff --git a/app/assets/database/queries/__init__.py b/app/assets/database/queries/__init__.py index 9949e84e1..38dd82d81 100644 --- a/app/assets/database/queries/__init__.py +++ b/app/assets/database/queries/__init__.py @@ -28,6 +28,7 @@ from app.assets.database.queries.asset_reference import ( get_reference_by_id, get_reference_with_owner_check, get_reference_ids_by_ids, + get_reference_paths_by_ids, get_references_by_paths_and_asset_ids, get_references_for_prefixes, get_unenriched_references, @@ -101,6 +102,7 @@ __all__ = [ "get_reference_by_id", "get_reference_with_owner_check", "get_reference_ids_by_ids", + "get_reference_paths_by_ids", "get_reference_tags", "get_references_by_paths_and_asset_ids", "get_references_for_prefixes", diff --git a/app/assets/database/queries/asset_reference.py b/app/assets/database/queries/asset_reference.py index 967b0e43a..13220270b 100644 --- a/app/assets/database/queries/asset_reference.py +++ b/app/assets/database/queries/asset_reference.py @@ -268,6 +268,8 @@ def list_references_page( order: str | None = None, after_cursor_value: object | None = None, after_cursor_id: str | None = None, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> tuple[list[AssetReference], dict[str, list[str]], int]: """List references with pagination, filtering, and sorting. @@ -293,7 +295,7 @@ def list_references_page( escaped, esc = escape_sql_like_string(name_contains) base = base.where(AssetReference.name.ilike(f"%{escaped}%", escape=esc)) - base = apply_tag_filters(base, include_tags, exclude_tags) + base = apply_tag_filters(base, include_tags, exclude_tags, any_tags) base = apply_metadata_filter(base, metadata_filter) sort = (sort or "created_at").lower() @@ -345,7 +347,7 @@ def list_references_page( count_stmt = count_stmt.where( AssetReference.name.ilike(f"%{escaped}%", escape=esc) ) - count_stmt = apply_tag_filters(count_stmt, include_tags, exclude_tags) + count_stmt = apply_tag_filters(count_stmt, include_tags, exclude_tags, any_tags) count_stmt = apply_metadata_filter(count_stmt, metadata_filter) total = int(session.execute(count_stmt).scalar_one() or 0) @@ -1062,6 +1064,27 @@ def get_references_by_paths_and_asset_ids( return winners +def get_reference_paths_by_ids( + session: Session, + reference_ids: list[str], +) -> dict[str, str]: + """Map reference id -> file_path for live, file-backed references.""" + if not reference_ids: + return {} + + paths: dict[str, str] = {} + for chunk in iter_chunks(reference_ids, MAX_BIND_PARAMS): + rows = session.execute( + select(AssetReference.id, AssetReference.file_path).where( + AssetReference.id.in_(chunk), + AssetReference.file_path.is_not(None), + AssetReference.deleted_at.is_(None), + ) + ) + paths.update({rid: fp for rid, fp in rows}) + return paths + + def get_reference_ids_by_ids( session: Session, reference_ids: list[str], diff --git a/app/assets/database/queries/common.py b/app/assets/database/queries/common.py index 89bb49327..7b0c211a0 100644 --- a/app/assets/database/queries/common.py +++ b/app/assets/database/queries/common.py @@ -60,10 +60,13 @@ def apply_tag_filters( stmt: sa.sql.Select, include_tags: Sequence[str] | None = None, exclude_tags: Sequence[str] | None = None, + any_tags: Sequence[str] | None = None, ) -> sa.sql.Select: - """include_tags: every tag must be present; exclude_tags: none may be present.""" + """include_tags: every tag must be present; any_tags: at least one must be + present; exclude_tags: none may be present.""" include_tags = normalize_tags(include_tags) exclude_tags = normalize_tags(exclude_tags) + any_tags = normalize_tags(any_tags) if include_tags: for tag_name in include_tags: @@ -74,6 +77,14 @@ def apply_tag_filters( ) ) + if any_tags: + stmt = stmt.where( + exists().where( + (AssetReferenceTag.asset_reference_id == AssetReference.id) + & (AssetReferenceTag.tag_name.in_(any_tags)) + ) + ) + if exclude_tags: stmt = stmt.where( ~exists().where( diff --git a/app/assets/database/queries/tags.py b/app/assets/database/queries/tags.py index 148f34801..e5f70e3df 100644 --- a/app/assets/database/queries/tags.py +++ b/app/assets/database/queries/tags.py @@ -340,6 +340,8 @@ def list_tag_counts_for_filtered_assets( name_contains: str | None = None, metadata_filter: dict | None = None, limit: int = 100, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> dict[str, int]: """Return tag counts for assets matching the given filters. @@ -359,7 +361,7 @@ def list_tag_counts_for_filtered_assets( escaped, esc = escape_sql_like_string(name_contains) ref_sq = ref_sq.where(AssetReference.name.ilike(f"%{escaped}%", escape=esc)) - ref_sq = apply_tag_filters(ref_sq, include_tags, exclude_tags) + ref_sq = apply_tag_filters(ref_sq, include_tags, exclude_tags, any_tags) ref_sq = apply_metadata_filter(ref_sq, metadata_filter) ref_sq = ref_sq.subquery() diff --git a/app/assets/scanner.py b/app/assets/scanner.py index 17a445d56..1791ebaf4 100644 --- a/app/assets/scanner.py +++ b/app/assets/scanner.py @@ -57,10 +57,11 @@ class _AssetAccumulator(TypedDict): refs: list[_RefInfo] +# Temp is deliberately absent: it is wiped before every scan, so walking it finds nothing. RootType = Literal["models", "input", "output"] -def get_prefixes_for_root(root: RootType) -> list[str]: +def get_scan_prefixes_for_root(root: RootType) -> list[str]: if root == "models": bases: list[str] = [] for _bucket, paths, _exts in get_comfy_models_folders(): @@ -73,10 +74,15 @@ def get_prefixes_for_root(root: RootType) -> list[str]: return [] -def get_all_known_prefixes() -> list[str]: - """Get all known asset prefixes across all root types.""" - all_roots: tuple[RootType, ...] = ("models", "input", "output") - return [p for root in all_roots for p in get_prefixes_for_root(root)] +def get_owned_prefixes() -> list[str]: + """Every directory an asset may live in; references outside these are marked missing.""" + scan_roots: tuple[RootType, ...] = ("models", "input", "output") + prefixes = [p for root in scan_roots for p in get_scan_prefixes_for_root(root)] + return prefixes + get_temp_prefixes() + + +def get_temp_prefixes() -> list[str]: + return [os.path.abspath(folder_paths.get_temp_directory())] def collect_models_files() -> list[str]: @@ -107,7 +113,21 @@ def sync_references_with_filesystem( collect_existing_paths: bool = False, update_missing_tags: bool = False, ) -> set[str] | None: - """Reconcile asset references with filesystem for a root. + return sync_prefixes_with_filesystem( + session, + get_scan_prefixes_for_root(root), + collect_existing_paths=collect_existing_paths, + update_missing_tags=update_missing_tags, + ) + + +def sync_prefixes_with_filesystem( + session, + prefixes: list[str], + collect_existing_paths: bool = False, + update_missing_tags: bool = False, +) -> set[str] | None: + """Reconcile asset references with filesystem under the given prefixes. - Toggle needs_verify per reference using mtime/size stat check - For hashed assets with at least one stat-unchanged ref: delete stale missing refs @@ -117,14 +137,13 @@ def sync_references_with_filesystem( Args: session: Database session - root: Root type to scan + prefixes: Absolute directory prefixes whose references to reconcile collect_existing_paths: If True, return set of surviving file paths update_missing_tags: If True, update 'missing' tags based on file status Returns: Set of surviving absolute paths if collect_existing_paths=True, else None """ - prefixes = get_prefixes_for_root(root) if not prefixes: return set() if collect_existing_paths else None @@ -251,6 +270,16 @@ def sync_root_safely(root: RootType) -> set[str]: return set() +def sync_temp_references_safely() -> None: + """Retire temp references whose file is gone; temp is never scanned, so nothing else stats them.""" + try: + with create_session() as sess: + sync_prefixes_with_filesystem(sess, get_temp_prefixes()) + sess.commit() + except Exception as e: + logging.exception("temp reference sync failed: %s", e) + + def mark_missing_outside_prefixes_safely(prefixes: list[str]) -> int: """Mark references as missing when outside the given prefixes. @@ -384,7 +413,7 @@ def get_unenriched_assets_for_roots( """ prefixes: list[str] = [] for root in roots: - prefixes.extend(get_prefixes_for_root(root)) + prefixes.extend(get_scan_prefixes_for_root(root)) if not prefixes: return [] diff --git a/app/assets/seeder.py b/app/assets/seeder.py index 2262928e5..134fc98a8 100644 --- a/app/assets/seeder.py +++ b/app/assets/seeder.py @@ -15,12 +15,13 @@ from app.assets.scanner import ( build_asset_specs, collect_paths_for_roots, enrich_assets_batch, - get_all_known_prefixes, - get_prefixes_for_root, + get_owned_prefixes, + get_scan_prefixes_for_root, get_unenriched_assets_for_roots, insert_asset_specs, mark_missing_outside_prefixes_safely, sync_root_safely, + sync_temp_references_safely, ) from app.database.db import dependencies_available @@ -413,7 +414,7 @@ class _AssetSeeder: ) return 0 - all_prefixes = get_all_known_prefixes() + all_prefixes = get_owned_prefixes() marked = mark_missing_outside_prefixes_safely(all_prefixes) if marked > 0: logging.info("Marked %d references as missing", marked) @@ -523,7 +524,7 @@ class _AssetSeeder: os.path.abspath(folder_paths.models_dir), ) else: - prefixes = get_prefixes_for_root(root) + prefixes = get_scan_prefixes_for_root(root) if prefixes: logging.info("Asset scan [%s] directories: %s", root, prefixes) @@ -548,10 +549,11 @@ class _AssetSeeder: return if self._prune_first: - all_prefixes = get_all_known_prefixes() + all_prefixes = get_owned_prefixes() marked = mark_missing_outside_prefixes_safely(all_prefixes) if marked > 0: logging.info("Marked %d refs as missing before scan", marked) + sync_temp_references_safely() if self._check_pause_and_cancel(): logging.info("Asset scan cancelled after pruning phase") diff --git a/app/assets/services/__init__.py b/app/assets/services/__init__.py index 03990966b..2747ad35b 100644 --- a/app/assets/services/__init__.py +++ b/app/assets/services/__init__.py @@ -4,6 +4,7 @@ from app.assets.services.asset_management import ( get_asset_by_hash, get_asset_detail, list_assets_page, + get_preview_file_paths, resolve_asset_for_download, set_asset_preview, update_asset_metadata, @@ -83,6 +84,7 @@ __all__ = [ "list_tags", "cleanup_unreferenced_assets", "remove_tags", + "get_preview_file_paths", "resolve_asset_for_download", "set_asset_preview", "update_asset_metadata", diff --git a/app/assets/services/asset_management.py b/app/assets/services/asset_management.py index a4c8b5a75..a04b7f21a 100644 --- a/app/assets/services/asset_management.py +++ b/app/assets/services/asset_management.py @@ -21,6 +21,7 @@ from app.assets.database.queries import ( reference_exists_for_asset_id, delete_reference_by_id, fetch_reference_and_asset, + get_reference_paths_by_ids, soft_delete_reference_by_id, fetch_reference_asset_and_tags, get_asset_by_hash as queries_get_asset_by_hash, @@ -279,6 +280,8 @@ def list_assets_page( sort: str = "created_at", order: str = "desc", after: str | None = None, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> ListAssetsResult: """List assets with optional cursor pagination. @@ -317,6 +320,7 @@ def list_assets_page( owner_id=owner_id, include_tags=include_tags, exclude_tags=exclude_tags, + any_tags=any_tags, name_contains=name_contains, metadata_filter=metadata_filter, limit=fetch_limit, @@ -421,6 +425,14 @@ def resolve_hash_to_path( ) +def get_preview_file_paths(preview_ids: list[str]) -> dict[str, str]: + """Map preview reference id -> file_path, in one query for the whole page.""" + if not preview_ids: + return {} + with create_session() as session: + return get_reference_paths_by_ids(session, reference_ids=preview_ids) + + def resolve_asset_for_download( reference_id: str, owner_id: str = "", diff --git a/app/assets/services/tagging.py b/app/assets/services/tagging.py index 5fa39d26a..69c8cf39c 100644 --- a/app/assets/services/tagging.py +++ b/app/assets/services/tagging.py @@ -85,6 +85,8 @@ def list_tag_histogram( name_contains: str | None = None, metadata_filter: dict | None = None, limit: int = 100, + # Appended last so pre-existing positional callers keep binding correctly. + any_tags: Sequence[str] | None = None, ) -> dict[str, int]: with create_session() as session: return list_tag_counts_for_filtered_assets( @@ -92,6 +94,7 @@ def list_tag_histogram( owner_id=owner_id, include_tags=include_tags, exclude_tags=exclude_tags, + any_tags=any_tags, name_contains=name_contains, metadata_filter=metadata_filter, limit=limit, diff --git a/blueprints/Image to Layers(Qwen-Image-Layered).json b/blueprints/Image to Layers(Qwen-Image-Layered).json index 7b44f0563..4b9bb8cf1 100644 --- a/blueprints/Image to Layers(Qwen-Image-Layered).json +++ b/blueprints/Image to Layers(Qwen-Image-Layered).json @@ -1,6 +1,6 @@ { "revision": 0, - "last_node_id": 176, + "last_node_id": 177, "last_link_id": 0, "nodes": [ { @@ -164,8 +164,8 @@ "version": 1, "state": { "lastGroupId": 8, - "lastNodeId": 176, - "lastLinkId": 380, + "lastNodeId": 177, + "lastLinkId": 381, "lastRerouteId": 0 }, "revision": 0, @@ -715,6 +715,88 @@ 1 ] }, + { + "id": 177, + "type": "LatentCut", + "pos": [ + 830, + -70 + ], + "size": [ + 270, + 170 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "localized_name": "samples", + "name": "samples", + "type": "LATENT", + "link": 142 + }, + { + "localized_name": "dim", + "name": "dim", + "type": "COMBO", + "widget": { + "name": "dim" + }, + "link": null + }, + { + "localized_name": "index", + "name": "index", + "type": "INT", + "widget": { + "name": "index" + }, + "link": null + }, + { + "localized_name": "amount", + "name": "amount", + "type": "INT", + "widget": { + "name": "amount" + }, + "link": null + } + ], + "outputs": [ + { + "localized_name": "LATENT", + "name": "LATENT", + "type": "LATENT", + "links": [ + 381 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.5.1", + "ue_properties": { + "widget_ue_connectable": {}, + "input_ue_unconnectable": {}, + "version": "7.7" + }, + "Node name for S&R": "LatentCut", + "enableTabs": false, + "tabWidth": 65, + "tabXOffset": 10, + "hasSecondTab": false, + "secondTabText": "Send Back", + "secondTabOffset": 80, + "secondTabWidth": 65 + }, + "widgets_values": [ + "t", + 1, + 16384 + ] + }, { "id": 76, "type": "LatentCutToBatch", @@ -734,7 +816,7 @@ "localized_name": "samples", "name": "samples", "type": "LATENT", - "link": 142 + "link": 381 }, { "localized_name": "dim", @@ -1434,7 +1516,7 @@ "id": 142, "origin_id": 3, "origin_slot": 0, - "target_id": 76, + "target_id": 177, "target_slot": 0, "type": "LATENT" }, @@ -1581,6 +1663,14 @@ "target_id": 39, "target_slot": 0, "type": "COMBO" + }, + { + "id": 381, + "origin_id": 177, + "origin_slot": 0, + "target_id": 76, + "target_slot": 0, + "type": "LATENT" } ], "extra": { diff --git a/comfy/background_removal/birefnet.py b/comfy/background_removal/birefnet.py index 78a80246e..ba3f710d4 100644 --- a/comfy/background_removal/birefnet.py +++ b/comfy/background_removal/birefnet.py @@ -433,19 +433,16 @@ class DeformableConv2d(nn.Module): def forward(self, x): offset = self.offset_conv(x) modulator = 2. * torch.sigmoid(self.modulator_conv(x)) - weight, bias, offload_info = comfy.ops.cast_bias_weight(self.regular_conv, x, offloadable=True) - - x = deform_conv2d( - input=x, - offset=offset, - weight=weight, - bias=None, - padding=self.padding, - mask=modulator, - stride=self.stride, - ) - comfy.ops.uncast_bias_weight(self.regular_conv, weight, bias, offload_info) - return x + with comfy.ops.CastBiasWeightContext(self.regular_conv, x, offloadable=True) as (weight, _bias): + return deform_conv2d( + input=x, + offset=offset, + weight=weight, + bias=None, + padding=self.padding, + mask=modulator, + stride=self.stride, + ) class BasicDecBlk(nn.Module): def __init__(self, in_channels=64, out_channels=64, inter_channels=64, device=None, dtype=None, operations=None): diff --git a/comfy/cli_args.py b/comfy/cli_args.py index ee9e1ce9f..659e772ed 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -74,7 +74,7 @@ parser.add_argument("--temp-directory", type=str, default=None, help="Set the Co parser.add_argument("--input-directory", type=str, default=None, help="Set the ComfyUI input directory. Overrides --base-directory.") parser.add_argument("--auto-launch", action="store_true", help="Automatically launch ComfyUI in the default browser.") parser.add_argument("--disable-auto-launch", action="store_true", help="Disable auto launching the browser.") -parser.add_argument("--cuda-device", type=str, default=None, metavar="DEVICE_ID", help="Set the ids of cuda devices this instance will use, as a comma-separated list (e.g. '0' or '0,1'). All other devices will not be visible.") +parser.add_argument("--cuda-device", type=str, default=None, metavar="DEVICE_ID", help="Set the ids of cuda devices this instance will use, as a comma-separated list (e.g. '0' or '0,1'), or 'all' to leave all currently visible devices available. All other devices will not be visible.") parser.add_argument("--default-device", type=int, default=None, metavar="DEFAULT_DEVICE_ID", help="Set the id of the default device, all other devices will stay visible.") cm_group = parser.add_mutually_exclusive_group() cm_group.add_argument("--cuda-malloc", action="store_true", help="Enable cudaMallocAsync (enabled by default for torch 2.0 and up).") @@ -149,6 +149,7 @@ attn_group.add_argument("--use-quad-cross-attention", action="store_true", help= attn_group.add_argument("--use-pytorch-cross-attention", action="store_true", help="Use the new pytorch 2.0 cross attention function.") attn_group.add_argument("--use-sage-attention", action="store_true", help="Use sage attention.") attn_group.add_argument("--use-flash-attention", action="store_true", help="Use FlashAttention.") +attn_group.add_argument("--use-ck-attention", action="store_true", help="Use Comfy Kitchen attention.") parser.add_argument("--disable-xformers", action="store_true", help="Disable xformers.") @@ -179,6 +180,7 @@ parser.add_argument("--disable-async-offload", action="store_true", help="Disabl parser.add_argument("--disable-dynamic-vram", action="store_true", help="Disable dynamic VRAM and use estimate based model loading.") parser.add_argument("--enable-dynamic-vram", action="store_true", help="Enable dynamic VRAM on systems where it's not enabled by default.") parser.add_argument("--fast-disk", action="store_true", help="Prefer disk-backed dynamic loading and offload over unpinned RAM. Can be faster for users with fast NVME disks.") +parser.add_argument("--disable-cuda-graphs", action="store_true", help="Disable CUDA graphs.") parser.add_argument("--force-non-blocking", action="store_true", help="Force ComfyUI to use non-blocking operations for all applicable tensors. This may improve performance on some non-Nvidia systems but can cause issues with some workflows.") diff --git a/comfy/clip_model.py b/comfy/clip_model.py index d7d3f994c..26cc5d7ee 100644 --- a/comfy/clip_model.py +++ b/comfy/clip_model.py @@ -314,13 +314,18 @@ class CLIPVisionModelProjection(torch.nn.Module): if "projection_dim" in config_dict: self.visual_projection = operations.Linear(config_dict["hidden_size"], config_dict["projection_dim"], bias=False) else: - self.visual_projection = lambda a: a + self.visual_projection = torch.nn.Identity() if "llava3" == config_dict.get("projector_type", None): self.multi_modal_projector = LlavaProjector(config_dict["hidden_size"], 4096, dtype, device, operations) else: self.multi_modal_projector = None + def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): + if "{}visual_projection.weight".format(prefix) not in state_dict: + self.visual_projection = torch.nn.Identity() + super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) + def forward(self, *args, **kwargs): x = self.vision_model(*args, **kwargs) out = self.visual_projection(x[2]) diff --git a/comfy/controlnet.py b/comfy/controlnet.py index 6dbbaa959..7e35fe027 100644 --- a/comfy/controlnet.py +++ b/comfy/controlnet.py @@ -381,13 +381,10 @@ class ControlLoraOps: self.bias = None def forward(self, input): - weight, bias, offload_stream = comfy.ops.cast_bias_weight(self, input, offloadable=True) - if self.up is not None: - x = torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias) - else: - x = torch.nn.functional.linear(input, weight, bias) - comfy.ops.uncast_bias_weight(self, weight, bias, offload_stream) - return x + with comfy.ops.CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + if self.up is None: + return torch.nn.functional.linear(input, weight, bias) + return torch.nn.functional.linear(input, weight + (torch.mm(self.up.flatten(start_dim=1), self.down.flatten(start_dim=1))).reshape(self.weight.shape).type(input.dtype), bias) class Conv2d(torch.nn.Module, comfy.ops.CastWeightBiasOp): def __init__( diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index c4270022b..34c6d700c 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -1,4 +1,5 @@ import torch +import comfy.nested_tensor class LatentFormat: scale_factor = 1.0 @@ -17,6 +18,9 @@ class LatentFormat: def process_out(self, latent): return latent / self.scale_factor + def fix_empty_latent(self, latent): + return latent + class SD15(LatentFormat): def __init__(self, scale_factor=0.18215): self.scale_factor = scale_factor @@ -573,6 +577,7 @@ class MiniMaxH3Video(LatentFormat): spacial_downscale_ratio = 16 temporal_downscale_ratio = 4 scale_factor = 1.0 + taesd_decoder_name = "taeh3" latent_rgb_factors = [ [-0.018555, 0.024344, -0.017536], @@ -606,6 +611,19 @@ class MiniMaxH3AV(MiniMaxH3Video): # max channels across the two streams (video 24, audio 32) so per-stream slices keep both streams whole latent_channels = 32 + def fix_empty_latent(self, latent): + video_latent_channels = MiniMaxH3Video.latent_channels + audio_latent_channels = 32 + audio_channels = 2 + frames_per_token = (1, 4, 4, 4, 4) + audio_frame_rescale = 5.0 / 3.0 + + video = latent[:, :video_latent_channels].clone() + frame_count = sum(frames_per_token[i % len(frames_per_token)] for i in range(video.shape[2])) + audio_t = round(frame_count * audio_frame_rescale) + audio = latent.new_zeros((latent.shape[0], audio_latent_channels, audio_channels, audio_t)) + return comfy.nested_tensor.NestedTensor((video, audio)) + class HunyuanVideo(LatentFormat): latent_channels = 16 latent_dimensions = 3 @@ -957,6 +975,11 @@ class ACEAudio15(LatentFormat): latent_dimensions = 1 temporal_downscale_ratio = 1764 +class MiniMaxMusic3(LatentFormat): + latent_channels = 128 + latent_dimensions = 1 + temporal_downscale_ratio = 512 + class ChromaRadiance(LatentFormat): latent_channels = 3 spacial_downscale_ratio = 1 diff --git a/comfy/ldm/lightricks/av_model.py b/comfy/ldm/lightricks/av_model.py index 8e360f6a8..c60148e2a 100644 --- a/comfy/ldm/lightricks/av_model.py +++ b/comfy/ldm/lightricks/av_model.py @@ -96,6 +96,8 @@ class BasicAVTransformerBlock(nn.Module): attn_precision=None, apply_gated_attention=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, dtype=None, device=None, operations=None, @@ -178,10 +180,10 @@ class BasicAVTransformerBlock(nn.Module): ) self.ff = FeedForward( - v_dim, dim_out=v_dim, glu=True, dtype=dtype, device=device, operations=operations + v_dim, dim_out=v_dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations ) self.audio_ff = FeedForward( - a_dim, dim_out=a_dim, glu=True, dtype=dtype, device=device, operations=operations + a_dim, dim_out=a_dim, glu=True, ff_bias=audio_ff_bias, dtype=dtype, device=device, operations=operations ) num_ada_params = ADALN_CROSS_ATTN_PARAMS_COUNT if cross_attention_adaln else ADALN_BASE_PARAMS_COUNT @@ -413,12 +415,16 @@ class LTXAVModel(LTXVModel): apply_gated_attention=False, caption_proj_before_connector=False, cross_attention_adaln=False, + ff_bias=True, + audio_ff_bias=True, + use_prompt_adaln_single=True, dtype=None, device=None, operations=None, **kwargs, ): # Store audio-specific parameters + self.audio_ff_bias = audio_ff_bias self.audio_in_channels = audio_in_channels self.audio_cross_attention_dim = audio_cross_attention_dim self.audio_attention_head_dim = audio_attention_head_dim @@ -451,6 +457,8 @@ class LTXAVModel(LTXVModel): timestep_scale_multiplier=timestep_scale_multiplier, caption_proj_before_connector=caption_proj_before_connector, cross_attention_adaln=cross_attention_adaln, + ff_bias=ff_bias, + use_prompt_adaln_single=use_prompt_adaln_single, dtype=dtype, device=device, operations=operations, @@ -475,7 +483,7 @@ class LTXAVModel(LTXVModel): operations=self.operations, ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.audio_prompt_adaln_single = AdaLayerNormSingle( self.audio_inner_dim, embedding_coefficient=2, @@ -606,6 +614,8 @@ class LTXAVModel(LTXVModel): a_context_dim=self.audio_cross_attention_dim, apply_gated_attention=self.apply_gated_attention, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, + audio_ff_bias=self.audio_ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -924,9 +934,15 @@ class LTXAVModel(LTXVModel): blocks_replace = patches_replace.get("dit", {}) prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.transformer_blocks), vx.device, transformer_options) + # Blocks whose self-attention should be perturbed to a value-passthrough (STG). + stg_self_attn_blocks = transformer_options.get("stg_self_attn_blocks", ()) + # Process transformer blocks for i, block in enumerate(self.transformer_blocks): comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, vx.device, block) + block_transformer_options = transformer_options + if i in stg_self_attn_blocks: + block_transformer_options = {**transformer_options, "stg_skip_self_attn": True} if ("double_block", i) in blocks_replace: def block_wrap(args): @@ -969,7 +985,7 @@ class LTXAVModel(LTXVModel): "a_cross_scale_shift_timestep": av_ca_audio_scale_shift_timestep, "v_cross_gate_timestep": av_ca_a2v_gate_noise_timestep, "a_cross_gate_timestep": av_ca_v2a_gate_noise_timestep, - "transformer_options": transformer_options, + "transformer_options": block_transformer_options, "self_attention_mask": self_attention_mask, "v_prompt_timestep": v_prompt_timestep, "a_prompt_timestep": a_prompt_timestep, @@ -993,7 +1009,7 @@ class LTXAVModel(LTXVModel): a_cross_scale_shift_timestep=av_ca_audio_scale_shift_timestep, v_cross_gate_timestep=av_ca_a2v_gate_noise_timestep, a_cross_gate_timestep=av_ca_v2a_gate_noise_timestep, - transformer_options=transformer_options, + transformer_options=block_transformer_options, self_attention_mask=self_attention_mask, v_prompt_timestep=v_prompt_timestep, a_prompt_timestep=a_prompt_timestep, diff --git a/comfy/ldm/lightricks/duration_head.py b/comfy/ldm/lightricks/duration_head.py new file mode 100644 index 000000000..7d45d1fa7 --- /dev/null +++ b/comfy/ldm/lightricks/duration_head.py @@ -0,0 +1,81 @@ +"""LTX 2.4 DurationHead: predicts the natural shot duration (in seconds) from +the caption connector token outputs, without running the diffusion pipeline. +""" + +import torch +import torch.nn.functional as F +from torch import nn + + +class AttentionPooler(nn.Module): + """Cross-attend ``num_queries`` learnable tokens against ``tokens``.""" + + def __init__(self, hidden_dim=256, num_queries=1, num_heads=4): + super().__init__() + self.num_queries = num_queries + self.query_tokens = nn.Parameter(torch.empty(num_queries, hidden_dim)) + self.cross_attn = nn.MultiheadAttention(embed_dim=hidden_dim, num_heads=num_heads, batch_first=True) + + def forward(self, tokens): + queries = self.query_tokens.unsqueeze(0).expand(tokens.shape[0], -1, -1) + pooled, _ = self.cross_attn(queries, tokens, tokens, need_weights=False) + return pooled + + +class DurationHead(nn.Module): + """Predict duration in seconds from one or both connector outputs.""" + + def __init__( + self, + video_cross_attention_dim=4096, + audio_cross_attention_dim=2048, + pooler_hidden_dim=256, + num_queries=1, + num_pooler_heads=4, + mlp_hidden=256, + ): + super().__init__() + self.video_input_proj = nn.Linear(video_cross_attention_dim, pooler_hidden_dim) + self.video_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.audio_input_proj = nn.Linear(audio_cross_attention_dim, pooler_hidden_dim) + self.audio_modality_emb = nn.Parameter(torch.empty(pooler_hidden_dim)) + self.attention_pooler = AttentionPooler( + hidden_dim=pooler_hidden_dim, num_queries=num_queries, num_heads=num_pooler_heads) + self.mlp_hidden = nn.Linear(pooler_hidden_dim * num_queries, mlp_hidden) + self.mlp_out = nn.Linear(mlp_hidden, 1) + + def forward(self, video_tokens=None, audio_tokens=None): + """``video_tokens``: (B, T_v, 4096), ``audio_tokens``: (B, T_a, 2048); + at least one required. Returns duration in seconds, shape (B,).""" + token_groups = [] + if video_tokens is not None: + token_groups.append(self.video_input_proj(video_tokens) + self.video_modality_emb) + if audio_tokens is not None: + token_groups.append(self.audio_input_proj(audio_tokens) + self.audio_modality_emb) + if not token_groups: + raise ValueError("DurationHead requires at least one of video_tokens / audio_tokens") + pooled = self.attention_pooler(torch.cat(token_groups, dim=1)) + pooled = pooled.reshape(pooled.shape[0], -1) + hidden = F.gelu(self.mlp_hidden(pooled), approximate="tanh") + return self.mlp_out(hidden).squeeze(-1).exp() + + +def normalize_state_dict(sd): + for prefix in ("model.diffusion_model.duration_head.", "duration_head."): + stripped = {k[len(prefix):]: v for k, v in sd.items() if k.startswith(prefix)} + if stripped: + return stripped + return sd + + +def seconds_to_num_frames(seconds, frame_rate, min_seconds, max_seconds, time_scale=8): + """Convert seconds to a frame count clamped to ``[min_seconds, max_seconds]`` + and snapped (floor) to the VAE's ``8k + 1`` causal temporal grid; snapping + that undershoots the minimum bumps up to the next grid point instead.""" + min_frames = max(1, round(min_seconds * frame_rate)) + max_frames = round(max_seconds * frame_rate) + raw_frames = max(min_frames, min(round(seconds * frame_rate), max_frames)) + frames = (raw_frames - 1) // time_scale * time_scale + 1 + if frames < min_frames: + frames = min(-(-(min_frames - 1) // time_scale) * time_scale + 1, max_frames) + return frames diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py index 1a6ddcc8d..9c412827f 100644 --- a/comfy/ldm/lightricks/embeddings_connector.py +++ b/comfy/ldm/lightricks/embeddings_connector.py @@ -50,6 +50,7 @@ class BasicTransformerBlock1D(nn.Module): context_dim=None, attn_precision=None, apply_gated_attention=False, + ff_bias=True, dtype=None, device=None, operations=None, @@ -74,6 +75,7 @@ class BasicTransformerBlock1D(nn.Module): dim, dim_out=dim, glu=True, + ff_bias=ff_bias, dtype=dtype, device=device, operations=operations, @@ -123,6 +125,7 @@ class Embeddings1DConnector(nn.Module): causal_temporal_positioning=False, num_learnable_registers: Optional[int] = 128, apply_gated_attention=False, + connector_ff_bias=True, dtype=None, device=None, operations=None, @@ -148,6 +151,7 @@ class Embeddings1DConnector(nn.Module): attention_head_dim, context_dim=cross_attention_dim, apply_gated_attention=apply_gated_attention, + ff_bias=connector_ff_bias, dtype=dtype, device=device, operations=operations, diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index f80bffba7..dcbfa43ad 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -303,22 +303,22 @@ class NormSingleLinearTextProjection(nn.Module): class GELU_approx(nn.Module): - def __init__(self, dim_in, dim_out, dtype=None, device=None, operations=None): + def __init__(self, dim_in, dim_out, bias=True, dtype=None, device=None, operations=None): super().__init__() - self.proj = operations.Linear(dim_in, dim_out, dtype=dtype, device=device) + self.proj = operations.Linear(dim_in, dim_out, bias=bias, dtype=dtype, device=device) def forward(self, x): return torch.nn.functional.gelu(self.proj(x), approximate="tanh") class FeedForward(nn.Module): - def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, dtype=None, device=None, operations=None): + def __init__(self, dim, dim_out, mult=4, glu=False, dropout=0.0, ff_bias=True, dtype=None, device=None, operations=None): super().__init__() inner_dim = int(dim * mult) - project_in = GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations) + project_in = GELU_approx(dim, inner_dim, bias=ff_bias, dtype=dtype, device=device, operations=operations) self.net = nn.Sequential( - project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, dtype=dtype, device=device) + project_in, nn.Dropout(dropout), operations.Linear(inner_dim, dim_out, bias=ff_bias, dtype=dtype, device=device) ) def forward(self, x): @@ -462,28 +462,34 @@ class CrossAttention(nn.Module): ) def forward(self, x, context=None, mask=None, pe=None, k_pe=None, transformer_options={}): + self_attn = context is None q = self.to_q(x) context = x if context is None else context k = self.to_k(context) v = self.to_v(context) - q = self.q_norm(q) - k = self.k_norm(k) - - # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent. - if pe is not None: - if k_pe is None and q.shape == k.shape: - q, k = apply_rotary_emb_qk(q, k, pe) - else: - q = apply_rotary_emb(q, pe) - k = apply_rotary_emb(k, pe if k_pe is None else k_pe) - - if mask is None: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) - elif isinstance(mask, GuideAttentionMask): - out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + # Spatio-Temporal Guidance (STG) perturbation: for the flagged self-attention + # layers, the attention degrades to a passthrough of the value projection (out = V). + if self_attn and transformer_options.get("stg_skip_self_attn", False): + out = v else: - out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + q = self.q_norm(q) + k = self.k_norm(k) + + # These norms span all heads, so the per-head RMS+RoPE kernel is not equivalent. + if pe is not None: + if k_pe is None and q.shape == k.shape: + q, k = apply_rotary_emb_qk(q, k, pe) + else: + q = apply_rotary_emb(q, pe) + k = apply_rotary_emb(k, pe if k_pe is None else k_pe) + + if mask is None: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, attn_precision=self.attn_precision, transformer_options=transformer_options) + elif isinstance(mask, GuideAttentionMask): + out = _attention_with_guide_mask(q, k, v, self.heads, mask, attn_precision=self.attn_precision, transformer_options=transformer_options) + else: + out = comfy.ldm.modules.attention.optimized_attention(q, k, v, self.heads, mask=mask, attn_precision=self.attn_precision, transformer_options=transformer_options) # Apply per-head gating if enabled if self.to_gate_logits is not None: @@ -502,7 +508,7 @@ ADALN_CROSS_ATTN_PARAMS_COUNT = 9 class BasicTransformerBlock(nn.Module): def __init__( - self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, dtype=None, device=None, operations=None + self, dim, n_heads, d_head, context_dim=None, attn_precision=None, cross_attention_adaln=False, ff_bias=True, dtype=None, device=None, operations=None ): super().__init__() @@ -518,7 +524,7 @@ class BasicTransformerBlock(nn.Module): device=device, operations=operations, ) - self.ff = FeedForward(dim, dim_out=dim, glu=True, dtype=dtype, device=device, operations=operations) + self.ff = FeedForward(dim, dim_out=dim, glu=True, ff_bias=ff_bias, dtype=dtype, device=device, operations=operations) self.attn2 = CrossAttention( query_dim=dim, @@ -717,6 +723,9 @@ class LTXBaseModel(torch.nn.Module, ABC): caption_proj_before_connector=False, cross_attention_adaln=False, caption_projection_first_linear=True, + ff_bias=True, + use_prompt_adaln_single=True, + use_keyframes_abs_pos_embedding=False, dtype=None, device=None, operations=None, @@ -746,6 +755,9 @@ class LTXBaseModel(torch.nn.Module, ABC): self.caption_proj_before_connector = caption_proj_before_connector self.cross_attention_adaln = cross_attention_adaln self.caption_projection_first_linear = caption_projection_first_linear + self.ff_bias = ff_bias + self.use_prompt_adaln_single = use_prompt_adaln_single + self.use_keyframes_abs_pos_embedding = use_keyframes_abs_pos_embedding # Common dimensions self.inner_dim = num_attention_heads * attention_head_dim @@ -773,12 +785,17 @@ class LTXBaseModel(torch.nn.Module, ABC): self.in_channels, self.inner_dim, bias=True, dtype=dtype, device=device ) + if self.use_keyframes_abs_pos_embedding: + self.keyframes_abs_pos_embedding = nn.Parameter(torch.zeros(1, self.inner_dim, dtype=dtype, device=device)) + else: + self.keyframes_abs_pos_embedding = None + embedding_coefficient = ADALN_CROSS_ATTN_PARAMS_COUNT if self.cross_attention_adaln else ADALN_BASE_PARAMS_COUNT self.adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=embedding_coefficient, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) - if self.cross_attention_adaln: + if self.cross_attention_adaln and self.use_prompt_adaln_single: self.prompt_adaln_single = AdaLayerNormSingle( self.inner_dim, embedding_coefficient=2, use_additional_conditions=False, dtype=dtype, device=device, operations=self.operations ) @@ -1070,6 +1087,7 @@ class LTXVModel(LTXBaseModel): self.attention_head_dim, context_dim=self.cross_attention_dim, cross_attention_adaln=self.cross_attention_adaln, + ff_bias=self.ff_bias, dtype=dtype, device=device, operations=self.operations, @@ -1099,6 +1117,15 @@ class LTXVModel(LTXBaseModel): grid_mask = None if keyframe_idxs is not None and keyframe_idxs.shape[2] > 0: + tokens_per_frame = self.tokens_per_latent_frame(additional_args["orig_shape"]) + if keyframe_idxs.shape[2] % tokens_per_frame != 0: + raise ValueError( + f"keyframe_idxs holds {keyframe_idxs.shape[2]} tokens, which is not a whole number of " + f"{tokens_per_frame}-token latent frames. The appended frames were recorded against a " + "different spatial resolution than the latent being sampled, so their positions would land " + "on the wrong tokens. Crop the guides and separate the generated keyframes before " + "upscaling the latent." + ) additional_args.update({ "orig_patchified_shape": list(x.shape)}) denoise_mask = self.patchifier.patchify(denoise_mask)[0] grid_mask = ~torch.any(denoise_mask < 0, dim=-1)[0] @@ -1141,8 +1168,64 @@ class LTXVModel(LTXBaseModel): additional_args["num_guide_tokens"] = keyframe_idxs.shape[2] x = self.patchify_proj(x) + x = self.apply_keyframes_abs_pos_embedding( + x, + pixel_coords, + orig_shape=additional_args["orig_shape"], + grid_mask=grid_mask, + num_guide_tokens=additional_args.get("num_guide_tokens", 0), + generated_keyframes=kwargs.get("generated_keyframes", None), + ) return x, pixel_coords, additional_args + def tokens_per_latent_frame(self, orig_shape): + """Token count of a single latent frame at the given latent shape.""" + patch_size = self.patchifier.patch_size + return (orig_shape[3] // patch_size[1]) * (orig_shape[4] // patch_size[2]) + + def keyframes_abs_pos_mask(self, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Per-token mask selecting the latents that encode a single standalone pixel frame. + + Returns a (batch, tokens) boolean mask over the already grid-filtered token sequence. + """ + temporal_start = pixel_coords[:, 0] + if temporal_start.ndim == 3: # (batch, tokens, [start, end]) + temporal_start = temporal_start[..., 0] + mask = temporal_start == 0 + if num_guide_tokens > 0: + mask[:, -num_guide_tokens:] = False + + if generated_keyframes is not None: + # The temporal patch size is always 1, so one latent frame is one row of tokens. + tokens_per_frame = self.tokens_per_latent_frame(orig_shape) + if generated_keyframes["tokens_per_frame"] != tokens_per_frame: + raise ValueError( + f"The generated keyframes were recorded at {generated_keyframes['tokens_per_frame']} tokens " + f"per latent frame but this latent has {tokens_per_frame}. Separate the generated keyframes " + "before upscaling the latent." + ) + first_token = generated_keyframes["first_latent_frame"] * tokens_per_frame + num_slot_tokens = generated_keyframes["num_keyframes"] * tokens_per_frame + slots = torch.zeros(orig_shape[2] * tokens_per_frame, dtype=torch.bool, device=mask.device) + slots[first_token:first_token + num_slot_tokens] = True + if grid_mask is not None: + slots = slots[grid_mask] + mask = mask | slots + + return mask + + def apply_keyframes_abs_pos_embedding(self, x, pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes): + """Add the learned keyframe marker to the single-pixel-frame tokens. + + A no-op for every checkpoint built without the parameter. + """ + if self.keyframes_abs_pos_embedding is None: + return x + + mask = self.keyframes_abs_pos_mask(pixel_coords, orig_shape, grid_mask, num_guide_tokens, generated_keyframes) + embedding = self.keyframes_abs_pos_embedding.to(device=x.device, dtype=x.dtype) + return x + mask.unsqueeze(-1).to(x.dtype) * embedding + def _build_guide_self_attention_mask(self, x, transformer_options, merged_args): """Build self-attention mask for per-guide attention attenuation. diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py index b4a8c7524..f5b1756d3 100644 --- a/comfy/ldm/lightricks/vae/audio_vae.py +++ b/comfy/ldm/lightricks/vae/audio_vae.py @@ -1,6 +1,5 @@ import json from dataclasses import dataclass -import math import torch import torchaudio @@ -186,7 +185,7 @@ class AudioVAE(torch.nn.Module): ) def num_of_latents_from_frames(self, frames_number: int, frame_rate: float) -> int: - return math.ceil((float(frames_number) / frame_rate) * self.latents_per_second) + return round((float(frames_number) / frame_rate) * self.latents_per_second) def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor: audio_channels = self.autoencoder.decoder.out_ch diff --git a/comfy/ldm/lightricks/vae/na_diffusion_decoder.py b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py new file mode 100644 index 000000000..ec539c665 --- /dev/null +++ b/comfy/ldm/lightricks/vae/na_diffusion_decoder.py @@ -0,0 +1,520 @@ +"""LTX 2.4 diffusion video VAE decoder (NADiffusionDecoder). + +Port of the reference ``DiffusionVideoDecoder`` without the NATTEN dependency: +``natten.na3d`` is replaced by ``comfy_kitchen.na3d``, which reproduces +NATTEN's semantics (window of exactly ``kernel_size`` per query, shifted +inward at grid boundaries, dilation 1) and dispatches cuda/triton/eager per +device and dtype (the eager backend covers CPU and fp32). + +Stages 1-4 deterministically upsample the latent into a context volume via +NA transformer blocks + linear pixel-shuffle upsamples. Stage 5 runs +``DiffusionNABlock``s that denoise patchified noised pixels ``x_t`` guided by +that context through AdaLN-Zero scale/shift. The 2.4 checkpoint is single-step +``x0``: one forward pass yields the pixels directly, no Euler loop. + +State dict keys match the shipped checkpoints directly (fused ``attn.qkv``, +``t_embedder.mlp.{0,2}``, ``shared_adaln.proj``); no rename pass is needed. +""" + +import math + +import torch +import torch.nn.functional as F +from einops import rearrange +from torch import nn +import comfy.model_management + +from comfy.ldm.lightricks.model import get_timestep_embedding +from .causal_video_autoencoder import Encoder, processor + +import comfy_kitchen + +# Token chunk for the SwiGLU MLP (bounds the [chunk, hidden] workspace). +MLP_TOKEN_CHUNK = 65536 + + +def rms_norm(x, weight, eps=1e-6): + if hasattr(F, "rms_norm"): + return F.rms_norm(x, (x.shape[-1],), weight=weight.to(x.dtype), eps=eps) + x_f = x.float() + x_f = x_f * torch.rsqrt(x_f.pow(2).mean(-1, keepdim=True) + eps) + return (x_f * weight.float()).to(x.dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim, eps=1e-6): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x): + return rms_norm(x, self.weight, self.eps) + + +def patchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b c (f p) (h q) (w r) -> b (c p r q) f h w", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +def unpatchify(x, patch_size_hw, patch_size_t=1): + if patch_size_hw == 1 and patch_size_t == 1: + return x + return rearrange(x, "b (c p r q) f h w -> b c (f p) (h q) (w r)", p=patch_size_t, q=patch_size_hw, r=patch_size_hw) + + +# --- Absolute per-axis RoPE (matches ltx-core rope.py numerics) --- + +def default_rope_dim_split(head_dim): + d_t = (head_dim // 4) // 2 * 2 + d_hw = (head_dim - d_t) // 2 + if d_hw % 2 != 0: + d_t -= 2 + d_hw = (head_dim - d_t) // 2 + return (d_t, d_hw, d_hw) + + +def rope_inv_freqs(dim, base=10000.0, device=None): + out_device = device + if not comfy.model_management.supports_fp64(device): + device = torch.device("cpu") + + exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim + return (1.0 / torch.pow(torch.tensor(float(base), dtype=torch.float64, device=device), exponents)).to(dtype=torch.float32, device=out_device) + + +def _rope_tables(lengths, inv_freqs, device): + """Precompute per-axis fp32 cos/sin tables for global 0-based positions.""" + tables = [] + for length, inv in zip(lengths, inv_freqs): + pos = torch.arange(length, dtype=torch.float32, device=device) + ang = pos[:, None] * inv[None, :] + tables.append((ang.cos(), ang.sin())) + return tables + + +def _rope_matrices_slice(tables, t0, t1, h, w): + """Per-token rotation matrices ``(1, ts*h*w, 1, hd/2, 2, 2)`` fp32 for + ``comfy_kitchen.rms_rope_`` (interleaved-pair convention), covering global + frames ``[t0, t1)`` of the axis-factorized tables.""" + parts = [] + for (c, s), sl in zip(tables, (slice(t0, t1), slice(None), slice(None))): + c, s = c[sl], s[sl] + parts.append(torch.stack([c, -s, s, c], dim=-1).reshape(c.shape[0], 1, 1, c.shape[1], 2, 2)) + ts = t1 - t0 + freqs = torch.cat([ + parts[0].expand(ts, h, w, -1, 2, 2), + parts[1].transpose(0, 1).expand(ts, h, w, -1, 2, 2), + parts[2].movedim(0, 2).expand(ts, h, w, -1, 2, 2), + ], dim=3) + return freqs.reshape(1, ts * h * w, 1, -1, 2, 2) + + +class NeighborhoodAttention3D(nn.Module): + """QKV (fused, matching checkpoint keys) + q/k RMSNorm + abs RoPE + NA.""" + + def __init__(self, dim, kernel_size, head_dim=64, rope_base=10000.0): + super().__init__() + self.dim = dim + self.num_heads = dim // head_dim + self.head_dim = head_dim + self.kernel_size = tuple(kernel_size) + self.scale = head_dim ** -0.5 + self.rope_split = default_rope_dim_split(head_dim) + self.rope_base = rope_base + + self.qkv = nn.Linear(dim, dim * 3, bias=True) + self.proj = nn.Linear(dim, dim, bias=True) + self.q_norm = RMSNorm(head_dim, eps=1e-6) + self.k_norm = RMSNorm(head_dim, eps=1e-6) + + def forward(self, x, pre=None, add_to=None): + """``pre`` (per-token norm/modulate) is applied slice-wise so the full + pre-attention tensor is never materialized; ``add_to`` streams the + output projection into it in place (residual add) and returns it. + Both bound peak memory without changing results.""" + batch, t, h, w, _ = x.shape + inv_freqs = tuple(rope_inv_freqs(d, self.rope_base, device=x.device) for d in self.rope_split) + tables = _rope_tables((t, h, w), inv_freqs, x.device) + shape = (batch, t, h, w, self.num_heads, self.head_dim) + q = torch.empty(shape, dtype=x.dtype, device=x.device) + k = torch.empty(shape, dtype=x.dtype, device=x.device) + v = torch.empty(shape, dtype=x.dtype, device=x.device) + q_weight = (self.q_norm.weight.detach() * self.scale).to(x.dtype) # scale commutes with the rotation + k_weight = self.k_norm.weight.detach().to(x.dtype) + chunk = max(1, (2 ** 25) // max(h * w * self.dim, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + qc, kc, vc = self.qkv(sl).chunk(3, dim=-1) + cshape = (batch, t1 - t0, h, w, self.num_heads, self.head_dim) + q[:, t0:t1] = qc.reshape(cshape) + k[:, t0:t1] = kc.reshape(cshape) + v[:, t0:t1] = vc.reshape(cshape) + freqs = _rope_matrices_slice(tables, t0, t1, h, w) + nt = (t1 - t0) * h * w + for b in range(batch): + comfy_kitchen.rms_rope_( + q[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + k[b, t0:t1].view(1, nt, self.num_heads, self.head_dim), + freqs, q_weight, k_weight) + out = comfy_kitchen.na3d(q, k, v, list(self.kernel_size), None, 1.0) + del q, k, v + out = out.reshape(batch, t, h, w, self.dim) + res = add_to if add_to is not None else torch.empty_like(out) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + if add_to is not None: + res[:, t0:t1] += self.proj(out[:, t0:t1]) + else: + res[:, t0:t1] = self.proj(out[:, t0:t1]) + return res + + +class SwiGLU(nn.Module): + """``w_down(silu(w_gate(x)) * w_up(x))``, chunked over tokens to bound the + ``[chunk, hidden]`` workspace.""" + + def __init__(self, dim, hidden_dim): + super().__init__() + self.w_up = nn.Linear(dim, hidden_dim, bias=False) + self.w_gate = nn.Linear(dim, hidden_dim, bias=False) + self.w_down = nn.Linear(hidden_dim, dim, bias=False) + + def forward(self, x, pre=None, add_to=None): + """``pre``/``add_to`` as in ``NeighborhoodAttention3D.forward``.""" + _, t, h, w, _ = x.shape + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + out = add_to if add_to is not None else torch.empty_like(x) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + sl = x[:, t0:t1] if pre is None else pre(x[:, t0:t1]) + y = self.w_down(F.silu(self.w_gate(sl)) * self.w_up(sl)) + if add_to is not None: + out[:, t0:t1] += y + else: + out[:, t0:t1] = y + return out + + +class NABlock(nn.Module): + """Pre-norm transformer block: NA -> SwiGLU MLP with residual adds.""" + + def __init__(self, dim, kernel_size, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x): + x = self.attn(x, pre=self.norm1, add_to=x) + return self.mlp(x, pre=self.norm2, add_to=x) + + +def modulate(x, scale, shift): + return x * (1.0 + scale) + shift + + +class AdaLNZero(nn.Module): + """``t_emb`` -> 7 (scale/shift/gate) chunks; gate slots unused (folded at export).""" + + NUM_CHUNKS = 7 + + def __init__(self, dim, t_emb_dim): + super().__init__() + self.proj = nn.Linear(t_emb_dim, self.NUM_CHUNKS * dim, bias=True) + + def forward(self, t_emb): + h = self.proj(F.silu(t_emb)) + return tuple(c[:, None, None, None, :] for c in h.chunk(self.NUM_CHUNKS, dim=-1)) + + +class DiffusionNABlock(nn.Module): + """NA + SwiGLU with shared AdaLN-Zero scale/shift (ungated residuals).""" + + def __init__(self, dim, kernel_size, context_channels, head_dim=64, mlp_ratio=4.0): + super().__init__() + self.context_proj = nn.Linear(context_channels, dim, bias=True) + self.scale_shift_table = nn.Parameter(torch.zeros(AdaLNZero.NUM_CHUNKS, dim)) + self.norm1 = RMSNorm(dim, eps=1e-6) + self.attn = NeighborhoodAttention3D(dim, kernel_size, head_dim=head_dim) + self.norm2 = RMSNorm(dim, eps=1e-6) + hidden = (int(dim * mlp_ratio) + 15) // 16 * 16 + self.mlp = SwiGLU(dim, hidden) + + def forward(self, x, latent_context, modulation): + scale_msa, shift_msa, _, scale_mlp, shift_mlp, _, _ = [ + modulation[i] + self.scale_shift_table[i].view(1, 1, 1, 1, -1) for i in range(AdaLNZero.NUM_CHUNKS) + ] + chunk = max(1, MLP_TOKEN_CHUNK // max(x.shape[2] * x.shape[3], 1)) + for t0 in range(0, x.shape[1], chunk): + x[:, t0:t0 + chunk] += self.context_proj(latent_context[:, t0:t0 + chunk]) + x = self.attn(x, pre=lambda s: modulate(self.norm1(s), scale_msa, shift_msa), add_to=x) + return self.mlp(x, pre=lambda s: modulate(self.norm2(s), scale_mlp, shift_mlp), add_to=x) + + +class LinearPixelShuffleUpsample(nn.Module): + """Linear channel-expand, then channels-last pixel shuffle.""" + + def __init__(self, in_channels, stride, out_channels_reduction_factor=1): + super().__init__() + self.stride = tuple(stride) + proj_out_channels = math.prod(stride) * in_channels // out_channels_reduction_factor + self.out_channels = proj_out_channels // math.prod(stride) + self.proj = nn.Linear(in_channels, proj_out_channels, bias=True) + + def forward(self, x, drop_leading_frame=True): + batch, t, h, w, _ = x.shape + p1, p2, p3 = self.stride + out = torch.empty((batch, t * p1, h * p2, w * p3, self.out_channels), dtype=x.dtype, device=x.device) + chunk = max(1, MLP_TOKEN_CHUNK // max(h * w, 1)) + for t0 in range(0, t, chunk): + t1 = min(t0 + chunk, t) + out[:, t0 * p1:t1 * p1] = rearrange( + self.proj(x[:, t0:t1]), "b t h w (c p1 p2 p3) -> b (t p1) (h p2) (w p3) c", + p1=p1, p2=p2, p3=p3, + ) + if p1 == 2 and drop_leading_frame: + # The causal temporal pixel-shuffle duplicates the leading frame. + out = out[:, 1:] + return out + + +class TimestepEmbedder(nn.Module): + """Sinusoidal(256) -> MLP. ``mlp.{0,2}`` naming matches the checkpoint.""" + + def __init__(self, t_emb_dim=384, freq_dim=256): + super().__init__() + self.freq_dim = freq_dim + self.mlp = nn.Sequential( + nn.Linear(freq_dim, t_emb_dim, bias=True), + nn.SiLU(), + nn.Linear(t_emb_dim, t_emb_dim, bias=True), + ) + + def forward(self, timestep, dtype): + emb = get_timestep_embedding(timestep.flatten(), self.freq_dim, flip_sin_to_cos=True, + downscale_freq_shift=0, scale=1) + return self.mlp(emb.to(dtype)) + + +class NADiffusionDecoder(nn.Module): + """Stages 1-4 (deterministic NA upsample) + stage-5 diffusion blocks. + + Input latent must already be un-normalized (the wrapper applies + ``per_channel_statistics.un_normalize``, same as the conv VAE path). + """ + + def __init__( + self, + in_channels=128, + out_channels=3, + patch_size=4, + head_dim=64, + stage_channels=(2048, 1024, 512, 512, 256), + stage_depths=(4, 6, 4, 2, 8), + stage_kernels=((3, 7, 7), (3, 7, 7), (3, 5, 5), (3, 5, 5), (11, 11, 11)), + upsamples=(((1, 2, 2), 2), ((2, 1, 1), 2), ((2, 2, 2), 1), ((2, 2, 2), 2)), + stage5_kernel=(11, 11, 11), + t_emb_dim=384, + default_num_inference_steps=1, + timestep_scale_multiplier=1000.0, + model_output_type="x0", + ): + super().__init__() + self.patch_size = patch_size + self.out_channels = out_channels + self.timestep_scale_multiplier = timestep_scale_multiplier + self.model_output_type = model_output_type + self.register_buffer( + "default_inference_timesteps", + torch.linspace(1.0, 1.0 / default_num_inference_steps, default_num_inference_steps), + persistent=False, + ) + self.temporal_upscale = math.prod(s[0] for s, _ in upsamples) + self.spatial_upscale = math.prod(s[1] for s, _ in upsamples) * patch_size + # NATTEN-style last-frame border mitigation: replicate the last latent + # frame through stages 1-4, crop the appendix off the context after. + self.trailing_pad_latent_frames = (stage_kernels[0][0] // 2) * 2 + + self.conv_in = nn.Linear(in_channels, stage_channels[0], bias=True) + + self.det_stages = nn.ModuleList() + self.upsamples = nn.ModuleList() + for stage_i in range(len(stage_channels) - 1): + c = stage_channels[stage_i] + self.det_stages.append(nn.ModuleList( + [NABlock(c, stage_kernels[stage_i], head_dim=head_dim) for _ in range(stage_depths[stage_i])] + )) + stride, reduction = upsamples[stage_i] + self.upsamples.append(LinearPixelShuffleUpsample(c, stride, out_channels_reduction_factor=reduction)) + + self.t_embedder = TimestepEmbedder(t_emb_dim=t_emb_dim) + + c5 = stage_channels[-1] + self.context_channels = c5 + noised_pixel_channels = out_channels * (patch_size ** 2) + self.conv_in_x_t = nn.Linear(noised_pixel_channels, c5, bias=True) + self.shared_adaln = AdaLNZero(c5, t_emb_dim) + self.diff_blocks = nn.ModuleList([ + DiffusionNABlock(c5, stage5_kernel, context_channels=c5, head_dim=head_dim) + for _ in range(stage_depths[-1]) + ]) + self.norm_out = RMSNorm(c5, eps=1e-6) + self.conv_out = nn.Linear(c5, noised_pixel_channels, bias=True) + + def forward_pre_diffusion(self, z, drop_leading_frame=True, pad_trailing=True): + """Stages 1-4: latent -> stage-5 context, channels-last. + + ``drop_leading_frame`` must be True only when ``z`` contains the + latent's true temporal origin (t=0); tiled callers decoding a later + temporal chunk pass False (the duplicate leading frame belongs solely + to the origin chunk). ``pad_trailing`` only for chunks containing the + latent's last frame.""" + n = self.trailing_pad_latent_frames if pad_trailing else 0 + if n > 0: + z = torch.cat([z, z[:, :, -1:].expand(-1, -1, n, -1, -1)], dim=2) + x = z.permute(0, 2, 3, 4, 1) + x = self.conv_in(x) + for stage_i, blocks in enumerate(self.det_stages): + for block in blocks: + x = block(x) + x = self.upsamples[stage_i](x, drop_leading_frame=drop_leading_frame) + if n > 0: + x = x[:, :-(n * self.temporal_upscale)] + return x + + def forward_diff_step(self, context, x_t, t): + x = patchify(x_t, patch_size_hw=self.patch_size, patch_size_t=1) + x = self.conv_in_x_t(x.permute(0, 2, 3, 4, 1)) + t_emb = self.t_embedder(self.timestep_scale_multiplier * t, dtype=x.dtype) + modulation = self.shared_adaln(t_emb) + for block in self.diff_blocks: + x = block(x, context, modulation) + x = self.norm_out(x) + x = self.conv_out(x) + x = x.permute(0, 4, 1, 2, 3) + return unpatchify(x, patch_size_hw=self.patch_size, patch_size_t=1) + + def forward(self, z, generator=None, drop_leading_frame=True, pad_trailing=True): + context = self.forward_pre_diffusion(z, drop_leading_frame=drop_leading_frame, pad_trailing=pad_trailing) + batch, t5, h5, w5, _ = context.shape + pixel_shape = (batch, self.out_channels, t5, h5 * self.patch_size, w5 * self.patch_size) + x_t = torch.randn(pixel_shape, dtype=z.dtype, device=z.device, generator=generator) + + timesteps = self.default_inference_timesteps.to(z.device) + num_steps = timesteps.shape[0] + for i in range(num_steps): + t_now = timesteps[i].expand(batch) + model_out = self.forward_diff_step(context, x_t, t_now) + if self.model_output_type == "x0": + x0 = model_out + if i == num_steps - 1: + return x0 + velocity = (x_t.float() - x0.float()) / timesteps[i] + else: # "v" + velocity = model_out.float() + if i == num_steps - 1: + return (x_t.float() - timesteps[i] * velocity).to(z.dtype) + t_next = timesteps[i + 1] if i + 1 < num_steps else torch.zeros_like(timesteps[i]) + x_t = (x_t.float() - (timesteps[i] - t_next) * velocity).to(z.dtype) + return x_t + + +LTX_24_VAE_CONFIG = { + "_class_name": "CausalDiffusionVAE", + "dims": 3, + "model_output_type": "x0", + "encoder": { + "dims": 3, + "in_channels": 3, + "out_channels": 128, + "blocks": [ + ["res_x", {"num_layers": 4}], + ["compress_space_res", {"multiplier": 2}], + ["res_x", {"num_layers": 6}], + ["compress_time_res", {"multiplier": 2}], + ["res_x", {"num_layers": 4}], + ["compress_all_res", {"multiplier": 2}], + ["res_x", {"num_layers": 2}], + ["compress_all_res", {"multiplier": 1}], + ["res_x", {"num_layers": 2}], + ], + "patch_size": 4, + "latent_log_var": "constant", + "norm_layer": "pixel_norm", + "base_channels": 128, + "spatial_padding_mode": "zeros", + }, + "decoder": { + "in_channels": 128, + "out_channels": 3, + "patch_size": 4, + "head_dim": 64, + "stage_channels": [2048, 1024, 512, 512, 256], + "stage_depths": [4, 6, 4, 2, 8], + "stage_kernels": [[3, 7, 7], [3, 7, 7], [3, 5, 5], [3, 5, 5], [11, 11, 11]], + "upsamples": [[[1, 2, 2], 2], [[2, 1, 1], 2], [[2, 2, 2], 1], [[2, 2, 2], 2]], + "stage5_kernel": [11, 11, 11], + "timestep_scale_multiplier": 1000.0, + "default_num_inference_steps": 1, + }, +} + + +class CausalDiffusionVAE(nn.Module): + """LTX 2.4 video VAE: conv encoder (shared with the 2.0 arch) + NA + diffusion decoder. Interface mirrors ``causal_video_autoencoder.VideoVAE``. + """ + + def __init__(self, config=None): + super().__init__() + if config is None: + config = LTX_24_VAE_CONFIG + self.config = config + enc = config.get("encoder", LTX_24_VAE_CONFIG["encoder"]) + dec = config.get("decoder", LTX_24_VAE_CONFIG["decoder"]) + dec_defaults = LTX_24_VAE_CONFIG["decoder"] + + self.encoder = Encoder( + dims=enc.get("dims", 3), + in_channels=enc.get("in_channels", 3), + out_channels=enc.get("out_channels", 128), + blocks=enc.get("blocks", LTX_24_VAE_CONFIG["encoder"]["blocks"]), + patch_size=enc.get("patch_size", 4), + latent_log_var=enc.get("latent_log_var", "constant"), + norm_layer=enc.get("norm_layer", "pixel_norm"), + spatial_padding_mode=enc.get("spatial_padding_mode", "zeros"), + base_channels=enc.get("base_channels", 128), + ) + + self.decoder = NADiffusionDecoder( + in_channels=dec.get("in_channels", 128), + out_channels=dec.get("out_channels", 3), + patch_size=dec.get("patch_size", 4), + head_dim=dec.get("head_dim", 64), + stage_channels=tuple(dec.get("stage_channels", dec_defaults["stage_channels"])), + stage_depths=tuple(dec.get("stage_depths", dec_defaults["stage_depths"])), + stage_kernels=tuple(tuple(k) for k in dec.get("stage_kernels", dec_defaults["stage_kernels"])), + upsamples=tuple((tuple(s), r) for s, r in dec.get("upsamples", dec_defaults["upsamples"])), + stage5_kernel=tuple(dec.get("stage5_kernel", dec_defaults["stage5_kernel"])), + t_emb_dim=dec.get("t_emb_dim", 384), + default_num_inference_steps=dec.get("default_num_inference_steps", 1), + timestep_scale_multiplier=dec.get("timestep_scale_multiplier", 1000.0), + model_output_type=config.get("model_output_type", "x0"), + ) + + self.per_channel_statistics = processor() + + def encode(self, x, device=None): + x = x[:, :, :max(1, 1 + ((x.shape[2] - 1) // 8) * 8), :, :] + means, logvar = torch.chunk(self.encoder(x, device=device), 2, dim=1) + return self.per_channel_statistics.normalize(means) + + def decode(self, x): + # Fixed-seed noise so decodes are reproducible TODO: expose? + generator = torch.Generator(device=x.device) + generator.manual_seed(0) + return self.decoder(self.per_channel_statistics.un_normalize(x), generator=generator) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index bc06288ab..765f21988 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -25,7 +25,7 @@ import comfy.model_prefetch import comfy.ops import comfy.patcher_extension import comfy.quant_ops -from comfy.ldm.modules.attention import optimized_attention +from comfy.ldm.modules.attention import AttentionTensorContainer, optimized_attention FRAME_PER_TOKEN = (1, 4, 4, 4, 4) FRAME_RESCALE = 5.0 / 3.0 @@ -74,6 +74,17 @@ def _axis_from_sqrt_area(dim, patch, sqrt_area): return (torch.arange(n, dtype=torch.float64) * (ratio / n) + (1.0 - ratio) / 2.0) * 32.0 +def mask_row_values(mask, latent_t, lat_h, lat_w): + # [T, H, W] denoise mask (1 = generate) -> per-2x2-patch-row float in [0, 1], + # None when every row fully generates + m = torch.nn.functional.pad(mask, (0, lat_w - mask.shape[-1], 0, lat_h - mask.shape[-2]), mode="replicate") + m = m.reshape(latent_t, lat_h // 2, 2, lat_w // 2, 2).amax(dim=(2, 4)) + values = m.reshape(-1) + if bool((values >= 1.0 - 1e-3).all()): + return None + return values + + def _frame_grid(h, w): # area-normalized (h, w) coordinates of one latent frame's 2x2-patch rows area = math.sqrt(h * w) @@ -91,6 +102,18 @@ def _video_t_grid(n, origin): return float(origin) + torch.cat([torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(0)]) +def _ref_t_span(blk): + # time-axis span a reference block occupies ahead of the target streams + kind = blk["kind"] + if kind == "image": + return 1.0 + if kind == "audio": + return float(blk["ref_audio_t"]) + if kind in ("video", "video_audio"): + return max(float(blk["ref_audio_t"]), sum(_video_t_spans(blk["latent_t"]))) + return 0.0 + + def _audio_grid(cursor, t, w_low, w_high): # channel-major stereo rows: t advances per latent frame, w pinned to the grid extremes per stereo channel, h stays 0 g = torch.zeros(t * 2, 3, dtype=torch.float64) @@ -165,9 +188,10 @@ class Attention(nn.Module): else: q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) - q = q.transpose(0, 1).unsqueeze(0) - k = k.transpose(0, 1).unsqueeze(0) - v = v.transpose(0, 1).unsqueeze(0) + v = v.clone() + q = AttentionTensorContainer(q.transpose(0, 1).unsqueeze(0)) + k = AttentionTensorContainer(k.transpose(0, 1).unsqueeze(0)) + v = AttentionTensorContainer(v.transpose(0, 1).unsqueeze(0)) out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) return self.out_proj(out.squeeze(0)) @@ -199,17 +223,22 @@ class AdalnProj(nn.Module): return x.chunk(self.expand, dim=-1) +def _mod_row(vecs, row, dtype): + # row is a mod-row index, or a per-token LongTensor of mod-row indices + return vecs[row].to(dtype) + + def _mod_scale_shift(h, shift, scale, segments): # segments: [(start, stop, mod_row)] covering h contiguously. for a, b, row in segments: - h[a:b].mul_(1.0 + scale[row].to(h.dtype)).add_(shift[row].to(h.dtype)) + h[a:b].mul_(1.0 + _mod_row(scale, row, h.dtype)).add_(_mod_row(shift, row, h.dtype)) return h def _mod_gate(x, gate, other, segments): # other is the fresh attn/mlp output: accumulate the gated residual into the stream in place, one fused kernel per segment for a, b, row in segments: - x[a:b].addcmul_(other[a:b], gate[row].to(x.dtype)) + x[a:b].addcmul_(other[a:b], _mod_row(gate, row, x.dtype)) return x @@ -275,19 +304,21 @@ class FinalLayer(nn.Module): self.audio_out = operations.Linear(hidden, audio_dim, bias=True, dtype=torch.float32, device=device) def forward(self, x, t_emb, video_seg, audio_seg): - # video_seg / audio_seg: (start, stop, timestep_row) of the target streams + # video_seg / audio_seg: (start, stop, row) of the target streams, where row + # is a mod-row index or a per-token blend (see _mod_row) shift, scale = self.adaln_proj(t_emb) - va, vb, vrow = video_seg - aa, ab, arow = audio_seg - hv = (self.norm(x[va:vb]) * (1.0 + scale[vrow]) + shift[vrow]).to(torch.float32) - ha = (self.norm(x[aa:ab]) * (1.0 + scale[arow]) + shift[arow]).to(torch.float32) - return self.video_out(hv), self.audio_out(ha) + + def mod(seg): + a, b, row = seg + return (self.norm(x[a:b]) * (1.0 + _mod_row(scale, row, scale.dtype)) + _mod_row(shift, row, shift.dtype)).to(torch.float32) + + return self.video_out(mod(video_seg)), self.audio_out(mod(audio_seg)) class PackedLayout: """Static packed-sequence structure for one shape/conditioning signature.""" - def __init__(self, text_len, latent_t, latent_h, latent_w, audio_t, keyframes=None, refs=None, frame_count=None): + def __init__(self, text_len, latent_t, latent_h, latent_w, audio_t, keyframes=None, refs=None): frame, w_grid = _frame_grid(latent_h, latent_w) frame_rows = frame.shape[0] @@ -298,29 +329,37 @@ class PackedLayout: img_pos, img_update = [], [] audio_pos, audio_update = [], [] - cursor = text_len row = text_len - if keyframes: - # fl2va: keyframe cond rows right after text, sharing the target spatial grid - for kf in keyframes: - pixel_index = kf["resolved_frame_index"] - if pixel_index == 0: - cond_t = float(text_len) - elif frame_count is not None and pixel_index == frame_count - 1: - cond_t = float(text_len) + sum(_video_t_spans(latent_t)) - FRAME_RESCALE - else: - raise ValueError("only first/last keyframe anchors are supported") - g = torch.empty(frame_rows, 3, dtype=torch.float64) - g[:, 0] = cond_t - g[:, 1:] = frame - segments.append(("cond", frame_rows)) - pos.append(g) - img_pos.append(torch.arange(row, row + frame_rows)) - img_update.append(torch.zeros(frame_rows, dtype=torch.bool)) - row += frame_rows - target_audio_w = (float(w_grid[0]), float(w_grid[-1])) + # refs pack between text and the targets, so the target timeline starts after their spans + cursor = float(text_len) + for blk in refs or (): + cursor += _ref_t_span(blk) + + if keyframes: + # fl2va: keyframe cond rows right after text, sharing the target spatial grid; + # anchors count from the target timeline origin, FRAME_RESCALE per pixel frame, 1.0 per audio latent frame + for kf in keyframes: + cond_t = cursor + FRAME_RESCALE * kf["resolved_frame_index"] + video_latent = kf.get("latent") + if video_latent is not None: + vt = video_latent.shape[2] + n = vt * frame_rows + segments.append(("cond", n)) + pos.append(_video_grid(vt, frame, cond_t)) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + audio_latent = kf.get("audio_latent") + if audio_latent is not None: + rt = audio_latent.shape[-1] + segments.append(("cond_audio", rt * 2)) + pos.append(_audio_grid(cond_t, rt, *target_audio_w)) + audio_pos.append(torch.arange(row, row + rt * 2)) + audio_update.append(torch.zeros(rt * 2, dtype=torch.bool)) + row += rt * 2 + if refs: cursor = float(text_len) for blk in refs: @@ -388,7 +427,7 @@ class PackedLayout: self.audio_update = torch.cat(audio_update) self.signature = (text_len, latent_t, latent_h, latent_w, audio_t) # contiguous segment table (start, stop, kind) - # kinds: text / cond / ref_img / ref_audio / audio / video + # kinds: text / cond / cond_audio / ref_img / ref_audio / audio / video # the packed sequence is uniform per segment in (modality tag, timestep class), # except the text span (tag runs resolved at forward time from the presentation tags) seg_abs = [] @@ -485,7 +524,7 @@ class MiniMaxH3Model(nn.Module): rows.append(r.to(device)) return torch.cat(rows, dim=0) if rows else None - def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs): + def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, denoise_mask=None, audio_denoise_mask=None, **kwargs): # the sampler carries the audio as (sigma_v / sigma_a) * x_audio; undo it outside # the wrappers so they and the network see the stream's own latent and velocity scale = float((minimax_payload or {}).get("audio_scale", 1.0)) @@ -502,7 +541,8 @@ class MiniMaxH3Model(nn.Module): self._forward, self, comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) - ).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload, **kwargs) + ).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload, + denoise_mask=denoise_mask, audio_denoise_mask=audio_denoise_mask, **kwargs) if scale != 1.0: # d/d(sigma_v) of the carried variable @@ -510,7 +550,7 @@ class MiniMaxH3Model(nn.Module): + (1.0 + (scale - 1.0) * sigma_a).to(out[1].dtype) * out[1]) return out - def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs): + def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, denoise_mask=None, audio_denoise_mask=None, **kwargs): video_x, audio_x = x[0], x[1] orig_t, orig_h, orig_w = video_x.shape[2], video_x.shape[3], video_x.shape[4] video_x = comfy.ldm.common_dit.pad_to_patch_size(video_x, self.patch_size) @@ -528,8 +568,7 @@ class MiniMaxH3Model(nn.Module): if layout is None or layout.signature != (text_len, latent_t, lat_h, lat_w, audio_t): layout = PackedLayout(text_len, latent_t, lat_h, lat_w, audio_t, keyframes=payload.get("keyframes"), - refs=payload.get("refs"), - frame_count=payload.get("frame_count")) + refs=payload.get("refs")) # model_base passes model_sampling.timestep(sigma) = sigma * 1000 shift_v = float(transformer_options.get("minimax_h3_sigma_shift_video", self.sigma_shift_video)) @@ -541,15 +580,46 @@ class MiniMaxH3Model(nn.Module): # distinct timesteps are known analytically: text/pad follow video, cond rows pin near 1 vis_aug = float(payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP)) aud_aug = float(payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP)) - has_vis_cond = any(k in ("cond", "ref_img") for _, _, k in layout.segments) - has_aud_cond = any(k == "ref_audio" for _, _, k in layout.segments) seg_t = {"text": t_v, "video": t_v, "audio": t_a, "cond": max(t_v, vis_aug), "ref_img": max(t_v, vis_aug), - "ref_audio": max(t_a, aud_aug)} - unique_t = sorted({t_v, t_a} | ({seg_t["cond"]} if has_vis_cond else set()) - | ({seg_t["ref_audio"]} if has_aud_cond else set())) + "cond_audio": max(t_a, aud_aug), "ref_audio": max(t_a, aud_aug)} + + # masked rows run at their own strength: mask value m puts a row at sigma = m * sigma_stream, + # so its label is 1 - m * sigma, clamped at the cond timestep for fully preserved rows + t_pin_v = max(t_v, VISUAL_COND_TIMESTEP) + t_pin_a = max(t_a, AUDIO_COND_TIMESTEP) + video_rows_t = None + audio_rows_t = None + if denoise_mask is not None: + m = mask_row_values(denoise_mask[0, 0].to(torch.float32), latent_t, lat_h, lat_w) + if m is not None: + rows_t = (1.0 - m * sigma_v.to(m.device)).clamp(max=t_pin_v) + if rows_t.unique().numel() == 1: + seg_t["video"] = float(rows_t[0]) + else: + video_rows_t = rows_t + if audio_denoise_mask is not None: + m = audio_denoise_mask[0, 0].to(torch.float32).reshape(-1) + if not bool((m >= 1.0 - 1e-3).all()): + sigma_a = 1.0 - t_a + rows_t = (1.0 - m * sigma_a).clamp(max=t_pin_a) + if rows_t.unique().numel() == 1: + seg_t["audio"] = float(rows_t[0]) + else: + audio_rows_t = rows_t + + unique_t = sorted({t_v, t_a} | {seg_t[k] for _, _, k in layout.segments} + | (set(video_rows_t.unique().tolist()) if video_rows_t is not None else set()) + | (set(audio_rows_t.unique().tolist()) if audio_rows_t is not None else set())) t_row = {t: i for i, t in enumerate(unique_t)} - seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "ref_audio": 2} + seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "cond_audio": 2, "ref_audio": 2} + + def rows_to_mod_index(rows_t, tag): + # per-row timestep values -> per-row mod-row indices into the t_emb table + levels = rows_t.unique() + base = torch.tensor([t_row[v] * 3 + tag for v in levels.tolist()], + dtype=torch.long, device=rows_t.device) + return base[torch.searchsorted(levels, rows_t)] text_tags = payload.get("text_token_tags") mod_segments = [] @@ -563,6 +633,10 @@ class MiniMaxH3Model(nn.Module): if i == b - a or tags[i] != tags[run_start]: mod_segments.append((a + run_start, a + i, row_base + int(tags[run_start]))) run_start = i + elif kind == "video" and video_rows_t is not None: + mod_segments.append((a, b, rows_to_mod_index(video_rows_t, seg_tag[kind]))) + elif kind == "audio" and audio_rows_t is not None: + mod_segments.append((a, b, rows_to_mod_index(audio_rows_t, seg_tag[kind]))) else: mod_segments.append((a, b, row_base + seg_tag[kind])) @@ -639,8 +713,16 @@ class MiniMaxH3Model(nn.Module): comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, None) # target streams are single contiguous segments (audio then video, last two) - video_seg = next((a, b, t_row[seg_t["video"]]) for a, b, k in layout.segments if k == "video") - audio_seg = next((a, b, t_row[seg_t["audio"]]) for a, b, k in layout.segments if k == "audio") + va, vb, _ = next(s for s in layout.segments if s[2] == "video") + aa, ab, _ = next(s for s in layout.segments if s[2] == "audio") + if video_rows_t is not None: + video_seg = (va, vb, rows_to_mod_index(video_rows_t, 0) // 3) + else: + video_seg = (va, vb, t_row[seg_t["video"]]) + if audio_rows_t is not None: + audio_seg = (aa, ab, rows_to_mod_index(audio_rows_t, 0) // 3) + else: + audio_seg = (aa, ab, t_row[seg_t["audio"]]) v, a = self.final_layer(h, t_emb, video_seg, audio_seg) video_out = unpatchify_video(v, latent_t, lat_h // 2, lat_w // 2, self.latents_dim, self.patch_size) diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py index 65d06f3e9..e1c146ec4 100644 --- a/comfy/ldm/minimax/vae.py +++ b/comfy/ldm/minimax/vae.py @@ -6,6 +6,7 @@ import torch import torch.nn as nn import torch.nn.functional as F +import comfy.model_management import comfy.ops import comfy.quant_ops import comfy.rmsnorm @@ -321,6 +322,8 @@ class ViT3DDecoder(nn.Module): # Full VAE class MiniMaxH3VideoVAE(nn.Module): + comfy_has_chunked_io = True + def __init__( self, in_channels=3, @@ -389,6 +392,23 @@ class MiniMaxH3VideoVAE(nn.Module): def _decode_pixels(self, z): return self.decoder(self.post_quant_conv(z)) + def _normalize_pixels(self, x): + return x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + + def _finalize_pixels(self, part): + # raw decoder output -> float32 pixels in [0, 1] (the VAE wrapper's process_output is identity) + part = part * self.pixel_std.to(device=part.device, dtype=torch.float32) + return part.add_(self.pixel_mean.to(device=part.device, dtype=torch.float32)).clamp_(0.0, 1.0) + + def decode_output_shape(self, input_shape): + b, c, t, h, w = input_shape + if t == 1: + frames = 1 + else: + pad_tokens, num_chunks = self._decode_temporal_chunks(t) + frames = self._decode_temporal_frame_plan(t + pad_tokens, num_chunks, pad_tokens) + return (b, self.decoder.out_channels, frames, h * self.vae_ratio, w * self.vae_ratio) + def _adaptive_encode(self, x): if self.tiling: return self.tiled_encode(x) @@ -521,18 +541,15 @@ class MiniMaxH3VideoVAE(nn.Module): # temporal chunking - def encode_temporal(self, x): - if x.shape[2] % self.clip_length != 0: - pad_size = (-x.shape[2]) % self.clip_length - pad_frames = x[:, :, -1:].repeat(1, 1, pad_size, 1, 1) - x = torch.cat([x, pad_frames], dim=2) - - num_chunks = x.shape[2] // self.clip_length - + def encode_temporal(self, x, device): + # chunked input io: x may live on the CPU, clips move to the device as they encode z_list = [] - for i in range(num_chunks): - clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :] - z_list.append(self._adaptive_encode(clip_x)) + for i in range(math.ceil(x.shape[2] / self.clip_length)): + clip_x = x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :].to(device) + if clip_x.shape[2] < self.clip_length: + pad_frames = clip_x[:, :, -1:].repeat(1, 1, self.clip_length - clip_x.shape[2], 1, 1) + clip_x = torch.cat([clip_x, pad_frames], dim=2) + z_list.append(self._adaptive_encode(self._normalize_pixels(clip_x))) z = torch.cat(z_list, dim=2) if self.token_drop > 0: @@ -577,43 +594,42 @@ class MiniMaxH3VideoVAE(nn.Module): total_frames += final_overlap_frames return total_frames - self._decode_temporal_pad_frames(z_len, pad_tokens) - def decode_temporal(self, z): - chunk_dec = self.tokens_chunk_size * self.vae_ratio_t - split_count = int(self.token_drop > 0) + 1 - - pseudo_total_tokens = z.shape[2] + self.token_drop - - pad_tokens = 0 - remainder = pseudo_total_tokens % self.tokens_chunk_size - if remainder != 0: - pad_tokens = self.tokens_chunk_size - remainder - pseudo_total_tokens += pad_tokens + def _decode_temporal_chunks(self, z_len): + pseudo_total_tokens = z_len + self.token_drop + pad_tokens = (-pseudo_total_tokens) % self.tokens_chunk_size + pseudo_total_tokens += pad_tokens num_chunks = pseudo_total_tokens // self.tokens_chunk_size - int(self.token_drop > 0) if num_chunks < 1: # too few tokens for one chunk (e.g. T_lat == 2): pad one extra chunk pad_tokens += self.tokens_chunk_size num_chunks += 1 + return pad_tokens, num_chunks + def decode_temporal(self, z, output_buffer=None): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + + if output_buffer is None: + # finalized chunks stream out of VRAM so the full video never sits on the GPU + output_buffer = torch.empty(self.decode_output_shape(z.shape), dtype=torch.float32, + device=comfy.model_management.intermediate_device()) + + pad_tokens, num_chunks = self._decode_temporal_chunks(z.shape[2]) if pad_tokens > 0: pad_z = z[:, :, -1:, :, :].repeat(1, 1, pad_tokens, 1, 1) z = torch.cat([z, pad_z], dim=2) - output_frames = self._decode_temporal_frame_plan(z.shape[2], num_chunks, pad_tokens) - - dec = None + dec = output_buffer dec_overlap = None write_pos = 0 def write_part(part): - nonlocal dec, write_pos + nonlocal write_pos part_frames = part.shape[2] if part_frames <= 0: return - if dec is None: - out_shape = list(part.shape) - out_shape[2] = output_frames - dec = torch.empty(out_shape, dtype=part.dtype, device=part.device) + part = self._finalize_pixels(part) copy_frames = min(part_frames, max(0, dec.shape[2] - write_pos)) if copy_frames > 0: dec[:, :, write_pos:write_pos + copy_frames, :, :].copy_( @@ -653,18 +669,18 @@ class MiniMaxH3VideoVAE(nn.Module): return dec - def encode(self, x): + def encode(self, x, device=None): # x: [B, 3, T, H, W] in [-1, 1] -> normalized latents [B, 24, T_lat, H/16, W/16] if x.ndim == 4: x = x.unsqueeze(2) - - x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + if device is None: + device = x.device if x.shape[2] == 1: - moments = self._adaptive_encode(x) + moments = self._adaptive_encode(self._normalize_pixels(x.to(device))) moments = moments[:, :, -1:, :, :] else: - moments = self.encode_temporal(x) + moments = self.encode_temporal(x, device) mean = torch.chunk(moments.float(), 2, dim=1)[0] @@ -679,18 +695,16 @@ class MiniMaxH3VideoVAE(nn.Module): def decode_tiled(self, z, **kwargs): return self.decode(z) - def decode(self, z): - # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> pixels [B, 3, T, H, W] in [-1, 1] + def decode(self, z, output_buffer=None): + # z: [B, 24, T_lat, H_lat, W_lat] normalized latents -> float32 pixels [B, 3, T, H, W] in [0, 1] latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(z) latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(z) z = z * latents_std + latents_mean if z.shape[2] == 1: - dec = self._adaptive_decode(z) - dec = dec[:, :, -1:, :, :] - else: - dec = self.decode_temporal(z) - - dec = dec.float() - dec.mul_(self.pixel_std.to(dec)).add_(self.pixel_mean.to(dec)).clamp_(0.0, 1.0).mul_(2.0).sub_(1.0) - return dec + dec = self._finalize_pixels(self._adaptive_decode(z)[:, :, -1:, :, :]) + if output_buffer is None: + return dec + output_buffer.copy_(dec) + return output_buffer + return self.decode_temporal(z, output_buffer) diff --git a/comfy/ldm/minimax_music/__init__.py b/comfy/ldm/minimax_music/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/comfy/ldm/minimax_music/ar.py b/comfy/ldm/minimax_music/ar.py new file mode 100644 index 000000000..2a8935318 --- /dev/null +++ b/comfy/ldm/minimax_music/ar.py @@ -0,0 +1,343 @@ +import dataclasses +import hashlib + +import torch +from torch import nn + +import comfy.model_management +import comfy.model_prefetch +import comfy.ops +import comfy.utils +from comfy.ldm.modules.attention import optimized_attention_for_device +from comfy.text_encoders.llama import Llama2_, Qwen3_8BConfig + +from .prompt import AUDIO_CODE_OFFSET, SPECIAL_TOKEN_IDS + + +CFG_SCALE = 1.5 +CFG_TOP_K = 50 +C0_VOCAB_SIZE = 16384 +MAX_PROMPT_TOKENS = 5000 +MAX_AUDIO_FRAMES = 9000 +AUDIO_FRAMES_PER_SECOND = 25 + + +def derive_seed(seed, *parts): + digest = hashlib.blake2b(digest_size=8, person=b"minimax-ttm") + digest.update(int(seed).to_bytes(8, "little", signed=False)) + for part in parts: + value = str(part).encode("utf-8") + digest.update(len(value).to_bytes(4, "little")) + digest.update(value) + return int.from_bytes(digest.digest(), "little") & ((1 << 63) - 1) + + +def sample_topk(logits, top_k, generator): + values = torch.nan_to_num(logits.float(), nan=-1e9, posinf=1e9, neginf=-1e9) + top_k = min(top_k, values.shape[-1]) + threshold = torch.topk(values, top_k, dim=-1).values[..., -1, None] + values = values.masked_fill(values < threshold, -float("inf")) + probabilities = torch.nan_to_num(torch.softmax(values, dim=-1), nan=0.0) + probabilities = probabilities / probabilities.sum(dim=-1, keepdim=True).clamp_min(1e-12) + return torch.multinomial(probabilities, 1, generator=generator).squeeze(-1) + + +class RVQAttention(nn.Module): + def __init__(self, hidden_size, num_heads, merged_qkv, dtype, device, operations): + super().__init__() + self.num_heads = num_heads + self.head_dim = hidden_size // num_heads + self.merged_qkv = merged_qkv + if merged_qkv: + self.qkv_proj = operations.Linear(hidden_size, hidden_size * 3, bias=False, dtype=dtype, device=device) + else: + self.q_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.k_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.v_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.o_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + + def forward(self, x): + batch, length, hidden_size = x.shape + if self.merged_qkv: + q, k, v = self.qkv_proj(x).chunk(3, dim=-1) + else: + q = self.q_proj(x) + k = self.k_proj(x) + v = self.v_proj(x) + q = q.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + k = k.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + v = v.reshape(batch, length, self.num_heads, self.head_dim).transpose(1, 2) + mask = torch.full((length, length), torch.finfo(q.dtype).min, device=q.device, dtype=q.dtype).triu_(1) + attention = optimized_attention_for_device(q.device, mask=True, small_input=True) + out = attention(q, k, v, self.num_heads, mask=mask, skip_reshape=True) + return self.o_proj(out) + + +class RVQRMSNorm(nn.Module): + def __init__(self, hidden_size, dtype, device): + super().__init__() + self.weight = nn.Parameter(torch.empty(hidden_size, dtype=dtype, device=device)) + + def forward(self, x): + return torch.nn.functional.rms_norm(x, (x.shape[-1],), comfy.ops.cast_to_input(self.weight, x), 1e-6) + + +class RVQMLP(nn.Module): + def __init__(self, hidden_size, intermediate_size, merged_mlp, dtype, device, operations): + super().__init__() + self.merged_mlp = merged_mlp + if merged_mlp: + self.gate_up_proj = operations.Linear(hidden_size, intermediate_size * 2, bias=False, dtype=dtype, device=device) + else: + self.gate_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.up_proj = operations.Linear(hidden_size, intermediate_size, bias=False, dtype=dtype, device=device) + self.down_proj = operations.Linear(intermediate_size, hidden_size, bias=False, dtype=dtype, device=device) + + def forward(self, x): + if self.merged_mlp: + return comfy.ops.linear_input_act(self.down_proj, self.gate_up_proj(x), "swiglu") + return self.down_proj(torch.nn.functional.silu(self.gate_proj(x)) * self.up_proj(x)) + + +class RVQDecoderBlock(nn.Module): + def __init__(self, hidden_size, num_heads, intermediate_size, merged_qkv, merged_mlp, dtype, device, operations): + super().__init__() + self.input_layernorm = RVQRMSNorm(hidden_size, dtype, device) + self.self_attn = RVQAttention(hidden_size, num_heads, merged_qkv, dtype, device, operations) + self.post_attention_layernorm = RVQRMSNorm(hidden_size, dtype, device) + self.mlp = RVQMLP(hidden_size, intermediate_size, merged_mlp, dtype, device, operations) + + def forward(self, x): + x = x + self.self_attn(self.input_layernorm(x)) + return x + self.mlp(self.post_attention_layernorm(x)) + + +class RVQDepthDecoder(nn.Module): + def __init__(self, config, dtype, device, operations): + super().__init__() + hidden_size = int(config["hidden_size"]) + audio_vocab_size = int(config["audio_vocab_size"]) + merged_qkv = config.get("decoder_merged_qkv", False) + merged_mlp = config.get("decoder_merged_mlp", False) + num_codebooks = int(config["audio_num_codebooks"]) + self.projection = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) + self.pos_embedding = operations.Embedding(16, hidden_size, dtype=dtype, device=device) + self.audio_heads = nn.ModuleList([ + operations.Linear(hidden_size, audio_vocab_size, bias=False, dtype=dtype, device=device) + for _ in range(num_codebooks - 1) + ]) + self.layers = nn.ModuleList([ + RVQDecoderBlock( + hidden_size, + int(config["decoder_num_heads"]), + int(config["decoder_intermediate_size"]), + merged_qkv, + merged_mlp, + dtype, + device, + operations, + ) + for _ in range(int(config["decoder_num_layers"])) + ]) + self.norm = RVQRMSNorm(hidden_size, dtype, device) + + def forward(self, sequence): + positions = torch.arange(sequence.shape[1], device=sequence.device) + x = sequence + self.pos_embedding(positions, out_dtype=sequence.dtype).unsqueeze(0) + for layer in self.layers: + x = layer(x) + return self.norm(x) + + +class MiniMaxMusic3AR(nn.Module): + def __init__(self, config, dtype, device, operations): + super().__init__() + config_fields = {field.name for field in dataclasses.fields(Qwen3_8BConfig)} + qwen_config = Qwen3_8BConfig(**{key: value for key, value in config.items() if key in config_fields}) + qwen_config.lm_head = False + qwen_config.fixed_kv = True + self.model = Llama2_(qwen_config, device=device, dtype=dtype, ops=operations) + self.model.prefetch_dynamic_vbars = True + self.model.graph_dynamic_vbar_blocks = True + self.model.lm_head = operations.Linear(qwen_config.hidden_size, qwen_config.vocab_size, bias=False, dtype=dtype, device=device) + self.model.lm_head_pruned = operations.Linear(qwen_config.hidden_size, C0_VOCAB_SIZE + 1, bias=False, dtype=dtype, device=device) + self.model.embed_tokens_prefill = operations.Embedding(AUDIO_CODE_OFFSET, qwen_config.hidden_size, dtype=dtype, device=device) + self.model.embed_tokens_audio = operations.Embedding(C0_VOCAB_SIZE, qwen_config.hidden_size, dtype=dtype, device=device) + self.model.pruned_lm_head = None + self.model.pruned_embedding = None + self.model.audio_extra_embedding = operations.Embedding( + int(config["audio_vocab_size"]) * (int(config["audio_num_codebooks"]) - 1), + qwen_config.hidden_size, + dtype=dtype, + device=device, + ) + self.model.audio_decoder = RVQDepthDecoder(config, dtype, device, operations) + self.audio_vocab_size = int(config["audio_vocab_size"]) + self.num_codebooks = int(config["audio_num_codebooks"]) + self.embedding_scale = self.num_codebooks ** -0.5 + + def _guided_c0(self, logits, cfg_scale, top_k): + conditioned = logits[0:1].float() + unconditioned = logits[1:2].float() + guided = unconditioned + (conditioned - unconditioned) * cfg_scale + threshold = torch.topk(conditioned, top_k, dim=-1).values[..., -1, None] + return guided.masked_fill(conditioned < threshold, -float("inf")) + + def _depth_codes(self, hidden, c0, c0_embed, generator, execution_dtype, cfg_scale, top_k): + decoder = self.model.audio_decoder + sequence = [decoder.projection(hidden).unsqueeze(1)] + sequence.append(decoder.projection(c0_embed).unsqueeze(1)) + codes = [c0] + hidden_parts = [] + for index in range(1, self.num_codebooks): + out = decoder(torch.cat(sequence, dim=1))[:, -1] + hidden_parts.append(out[:1].detach()) + logits = decoder.audio_heads[index - 1](out) + conditioned = logits[:1].float() + unconditioned = logits[1:2].float() + code = sample_topk(unconditioned + (conditioned - unconditioned) * cfg_scale, top_k, generator).repeat(2) + codes.append(code) + if index < self.num_codebooks - 1: + embedding = self.model.audio_extra_embedding( + code + (index - 1) * self.audio_vocab_size, + out_dtype=execution_dtype, + ) + sequence.append(decoder.projection(embedding).unsqueeze(1)) + return torch.stack(codes, dim=1), torch.cat(hidden_parts, dim=-1) + + def _embed_c0(self, codes, execution_dtype): + if self.model.pruned_embedding: + return self.model.embed_tokens_audio(codes, out_dtype=execution_dtype) + return self.model.embed_tokens(codes + AUDIO_CODE_OFFSET, out_dtype=execution_dtype) + + def _embed_audio_frame(self, codes, execution_dtype): + c0 = self._embed_c0(codes[:, 0], execution_dtype) + offsets = torch.arange(self.num_codebooks - 1, device=codes.device) * self.audio_vocab_size + extra = self.model.audio_extra_embedding(codes[:, 1:] + offsets.unsqueeze(0), out_dtype=execution_dtype).sum(dim=1) + return ((c0 + extra) * self.embedding_scale).unsqueeze(1) + + def _sample_c0(self, hidden, cfg_scale, top_k, generator, vocab_mask): + if self.model.pruned_lm_head: + guided = self._guided_c0(self.model.lm_head_pruned(hidden).float(), cfg_scale, top_k) + code = sample_topk(guided, top_k, generator) + stop_token = 0 + offset = 1 + else: + logits = self.model.lm_head(hidden).float() + stop_token = SPECIAL_TOKEN_IDS["<|audio_end|>"] + logits = logits.masked_fill(vocab_mask, -float("inf")) + guided = self._guided_c0(logits, cfg_scale, top_k).masked_fill(vocab_mask, -float("inf")) + code = sample_topk(guided, top_k, generator) + offset = AUDIO_CODE_OFFSET + return torch.where(code == stop_token, 0, code - offset), code, stop_token + + def generate(self, input_ids, seed, max_audio_frames, device, cfg_scale=CFG_SCALE, top_k=CFG_TOP_K): + prompt_tokens = int(input_ids.shape[1]) + if prompt_tokens > MAX_PROMPT_TOKENS: + raise ValueError(f"MiniMax Music3 prompt has {prompt_tokens} tokens; maximum is {MAX_PROMPT_TOKENS}") + + input_ids = input_ids.to(device) + if comfy.model_management.should_use_bf16(device): + execution_dtype = torch.bfloat16 + else: + execution_dtype = torch.float32 + unconditioned = input_ids.clone() + unconditioned[:, 1:-2] = SPECIAL_TOKEN_IDS["<|audio_cfg|>"] + text_ids = torch.cat((input_ids, unconditioned), dim=0) + if self.model.pruned_embedding: + text_embeds = self.model.embed_tokens_prefill(text_ids, out_dtype=execution_dtype) + else: + text_embeds = self.model.embed_tokens(text_ids, out_dtype=execution_dtype) + decode_limit = min(int(max_audio_frames), MAX_AUDIO_FRAMES) + past = self.model.init_kv_cache(2, prompt_tokens + decode_limit + 1, device, execution_dtype) + output = self.model(None, embeds=text_embeds, past_key_values=past, dtype=execution_dtype) + last_hidden = output[0][:, -1] + past = output[2] + + generator = torch.Generator(device=device).manual_seed(derive_seed(seed, "ar")) + decoder = self.model.audio_decoder + depth_io = { + "hidden": torch.empty_like(last_hidden), + "c0": torch.empty((last_hidden.shape[0],), dtype=torch.long, device=device), + "c0_embed": torch.empty_like(last_hidden), + "codes": torch.empty((last_hidden.shape[0], self.num_codebooks), dtype=torch.long, device=device), + "depth_hidden": torch.empty((1, last_hidden.shape[-1] * (self.num_codebooks - 1)), dtype=execution_dtype, device=device), + } + decoder._comfy_cross_step_state = depth_io + comfy.model_management._register_cross_step(decoder) + hidden_frames = [] + pending_code = None + stop_token = None + pending_event = None + pending_hidden = None + progress = comfy.utils.ProgressBar(decode_limit) + cuda_device = torch.device(device).type == "cuda" + vocab_mask = None + if not self.model.pruned_lm_head: + vocab_mask = torch.ones(self.model.vocab_size, dtype=torch.bool, device=device) + vocab_mask[AUDIO_CODE_OFFSET:AUDIO_CODE_OFFSET + C0_VOCAB_SIZE] = False + vocab_mask[SPECIAL_TOKEN_IDS["<|audio_end|>"]] = False + + for frame_index in comfy.utils.model_trange(decode_limit + 1, desc="AR sampling"): + comfy.model_management.throw_exception_if_processing_interrupted() + if pending_code is not None: + if pending_event is not None: + pending_event.synchronize() + if int(pending_code.item()) == stop_token: + pending_hidden = None + break + if pending_hidden is not None: + hidden_frames.append(pending_hidden) + progress.update_absolute(len(hidden_frames)) + if len(hidden_frames) >= decode_limit: + break + + c0, code_or_stop, stop_token = self._sample_c0(last_hidden, cfg_scale, top_k, generator, vocab_mask) + if pending_code is None: + pending_code = torch.empty_like(code_or_stop, device="cpu", pin_memory=cuda_device) + if cuda_device: + pending_event = torch.cuda.Event() + pending_code.copy_(code_or_stop, non_blocking=cuda_device) + if pending_event is not None: + pending_event.record() + + c0 = c0.repeat(2) + c0_embed = self._embed_c0(c0, execution_dtype) + depth_io["hidden"].copy_(last_hidden) + depth_io["c0"].copy_(c0) + depth_io["c0_embed"].copy_(c0_embed) + + def depth_core(): + codes, depth_hidden = self._depth_codes( + depth_io["hidden"], depth_io["c0"], depth_io["c0_embed"], generator, execution_dtype, cfg_scale, top_k + ) + depth_io["codes"].copy_(codes) + depth_io["depth_hidden"].copy_(depth_hidden) + + depth_queue = comfy.model_prefetch.make_prefetch_queue( + [[decoder, self.model.audio_extra_embedding]], device, {"prefetch_dynamic_vbars": True} + ) + comfy.model_prefetch.prefetch_queue_pop( + depth_queue, device, decoder, execution_dtype, core=depth_core, enable_graph=True, generator=generator + ) + comfy.model_prefetch.prefetch_queue_pop(depth_queue, device, None) + feedback_codes = depth_io["codes"] + depth_hidden = depth_io["depth_hidden"] + frame_hidden = torch.cat((last_hidden[:1].detach(), depth_hidden), dim=-1) + if frame_index > 0: + pending_hidden = frame_hidden[0].clone() + + feedback = self._embed_audio_frame(feedback_codes, execution_dtype) + output = self.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype) + last_hidden = output[0][:, -1] + past = output[2] + + if pending_hidden is not None and len(hidden_frames) < decode_limit: + if pending_event is not None: + pending_event.synchronize() + if int(pending_code.item()) != stop_token: + hidden_frames.append(pending_hidden) + + if not hidden_frames: + raise ValueError("MiniMax Music3 generated zero audio frames") + return torch.stack(hidden_frames).to(device="cpu") diff --git a/comfy/ldm/minimax_music/dav.py b/comfy/ldm/minimax_music/dav.py new file mode 100644 index 000000000..d442559f4 --- /dev/null +++ b/comfy/ldm/minimax_music/dav.py @@ -0,0 +1,137 @@ +import math + +import torch +from torch import nn + +import comfy.ops + + +def snake(x, alpha): + shape = x.shape + flat = x.reshape(shape[0], shape[1], -1) + alpha = comfy.ops.cast_to_input(alpha, flat) + flat = flat + (alpha + 1e-9).reciprocal() * torch.sin(alpha * flat).pow(2) + return flat.reshape(shape) + + +class Snake1d(nn.Module): + def __init__(self, channels, dtype, device): + super().__init__() + self.alpha = nn.Parameter(torch.empty(1, channels, 1, dtype=dtype, device=device)) + + def forward(self, x): + return snake(x, self.alpha) + + +def _weight_norm_conv(operations, *args, **kwargs): + return nn.utils.parametrizations.weight_norm(operations.Conv1d(*args, **kwargs)) + + +def _weight_norm_conv_transpose(operations, *args, **kwargs): + return nn.utils.parametrizations.weight_norm(operations.ConvTranspose1d(*args, **kwargs)) + + +class ResidualUnit(nn.Module): + def __init__(self, dim, dilation, dtype, device, operations): + super().__init__() + padding = 3 * dilation + self.block = nn.Sequential( + Snake1d(dim, dtype, device), + _weight_norm_conv( + operations, + dim, + dim, + kernel_size=7, + dilation=dilation, + padding=padding, + dtype=dtype, + device=device, + ), + Snake1d(dim, dtype, device), + _weight_norm_conv(operations, dim, dim, kernel_size=1, dtype=dtype, device=device), + ) + + def forward(self, x): + residual = self.block(x) + if residual.shape[-1] != x.shape[-1]: + padding = (x.shape[-1] - residual.shape[-1]) // 2 + x = x[..., padding:x.shape[-1] - padding] + return x + residual + + +class DecoderBlock(nn.Module): + def __init__(self, input_dim, output_dim, stride, dtype, device, operations): + super().__init__() + self.block = nn.Sequential( + Snake1d(input_dim, dtype, device), + _weight_norm_conv_transpose( + operations, + input_dim, + output_dim, + kernel_size=2 * stride, + stride=stride, + padding=math.ceil(stride / 2), + dtype=dtype, + device=device, + ), + ResidualUnit(output_dim, 1, dtype, device, operations), + ResidualUnit(output_dim, 3, dtype, device, operations), + ResidualUnit(output_dim, 9, dtype, device, operations), + ) + + def forward(self, x): + return self.block(x) + + +class Decoder(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + layers = [ + _weight_norm_conv( + operations, + 1024, + 1536, + kernel_size=7, + padding=3, + dtype=dtype, + device=device, + ) + ] + channels = 1536 + output_dim = channels + for index, stride in enumerate((8, 8, 4, 2)): + input_dim = channels // (2 ** index) + output_dim = channels // (2 ** (index + 1)) + layers.append(DecoderBlock(input_dim, output_dim, stride, dtype, device, operations)) + layers.extend(( + Snake1d(output_dim, dtype, device), + _weight_norm_conv( + operations, + output_dim, + 1, + kernel_size=7, + padding=3, + dtype=dtype, + device=device, + ), + nn.Tanh(), + )) + self.model = nn.Sequential(*layers) + + def forward(self, x): + return self.model(x) + + +class MiniMaxMusic3DAV(nn.Module): + def __init__(self, dtype=None, device=None, operations=None): + super().__init__() + self.dec_in_proj = operations.Conv1d(64, 1024, kernel_size=1, dtype=dtype, device=device) + self.decoder = Decoder(dtype, device, operations) + + def decode(self, latent): + batch, _, frames = latent.shape + folded = latent.reshape(batch * 2, 64, frames) + waveform = self.decoder(self.dec_in_proj(folded)) + return waveform.reshape(batch, 2, -1) + + forward = decode diff --git a/comfy/ldm/minimax_music/dit.py b/comfy/ldm/minimax_music/dit.py new file mode 100644 index 000000000..211e0d7db --- /dev/null +++ b/comfy/ldm/minimax_music/dit.py @@ -0,0 +1,213 @@ +import math + +import torch +from torch import nn + +import comfy.model_management +import comfy.ops +import comfy.quant_ops +from comfy.ldm.modules.attention import optimized_attention_for_device + + +MAX_CONDITION_FRAMES = 200 +CONDITION_HOP_FRAMES = 100 + + +def latent_length(audio_frames): + return max(1, int(audio_frames * 44100 / 24000 * 960 / 512)) + + +class FourierFeatures(nn.Module): + def __init__(self, in_features, out_features, dtype, device): + super().__init__() + self.weight = nn.Parameter(torch.empty(out_features // 2, in_features, dtype=dtype, device=device)) + + def forward(self, value): + weight = comfy.ops.cast_to_input(self.weight, value) + features = 2.0 * math.pi * value @ weight.T + return torch.cat((features.cos(), features.sin()), dim=-1) + + +class LayerNorm(nn.Module): + def __init__(self, dim, dtype, device): + super().__init__() + self.gamma = nn.Parameter(torch.empty(dim, dtype=dtype, device=device)) + self.register_buffer("beta", torch.empty(dim, dtype=dtype, device=device)) + + def forward(self, x): + return torch.nn.functional.layer_norm( + x, + (x.shape[-1],), + comfy.ops.cast_to_input(self.gamma, x), + comfy.ops.cast_to_input(self.beta, x), + ) + + +class RotaryEmbedding(nn.Module): + def __init__(self, dim, dtype, device): + super().__init__() + self.register_buffer("inv_freq", torch.empty(dim // 2, dtype=dtype, device=device)) + + def forward_from_seq_len(self, length, device, dtype): + positions = torch.arange(length, device=device, dtype=torch.float32) + frequencies = torch.outer(positions, comfy.ops.cast_to_input(self.inv_freq, positions)) + frequencies = frequencies.to(dtype) + cos, sin = frequencies.cos(), frequencies.sin() + return torch.stack((cos, -sin, sin, cos), dim=-1).reshape(1, 1, length, frequencies.shape[-1], 2, 2) + + +def _apply_rope(x, rotation_matrix): + x_dtype = x.dtype + x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(rotation_matrix.dtype) + x = rotation_matrix[..., 0] * x[..., 0] + rotation_matrix[..., 1] * x[..., 1] + return x.movedim(-1, -2).flatten(-2).to(x_dtype) + + +class Attention(nn.Module): + def __init__(self, dim, dim_heads, dtype, device, operations): + super().__init__() + self.num_heads = dim // dim_heads + self.dim_heads = dim_heads + self.to_qkv = operations.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device) + self.to_out = operations.Linear(dim, dim, bias=False, dtype=dtype, device=device) + + def forward(self, x, rotation_matrix): + batch, length, dim = x.shape + q, k, v = self.to_qkv(x).chunk(3, dim=-1) + q = q.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + k = k.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + v = v.reshape(batch, length, self.num_heads, self.dim_heads).transpose(1, 2) + rotary_dims = rotation_matrix.shape[-3] * 2 + if comfy.model_management.in_training: + q = torch.cat((_apply_rope(q[..., :rotary_dims], rotation_matrix), q[..., rotary_dims:]), dim=-1) + k = torch.cat((_apply_rope(k[..., :rotary_dims], rotation_matrix), k[..., rotary_dims:]), dim=-1) + else: + rotated_q, rotated_k = comfy.quant_ops.ck.apply_rope_split_half(q[..., :rotary_dims], k[..., :rotary_dims], rotation_matrix) + q = torch.cat((rotated_q, q[..., rotary_dims:]), dim=-1) + k = torch.cat((rotated_k, k[..., rotary_dims:]), dim=-1) + attention = optimized_attention_for_device(q.device) + out = attention(q, k, v, self.num_heads, skip_reshape=True) + return self.to_out(out) + + +class GLU(nn.Module): + def __init__(self, dim, inner_dim, dtype, device, operations): + super().__init__() + self.proj = operations.Linear(dim, inner_dim * 2, dtype=dtype, device=device) + + def forward(self, x): + value, gate = self.proj(x).chunk(2, dim=-1) + return value * torch.nn.functional.silu(gate) + + +class FeedForward(nn.Module): + def __init__(self, dim, inner_dim, dtype, device, operations): + super().__init__() + self.ff = nn.Sequential( + GLU(dim, inner_dim, dtype, device, operations), + nn.Identity(), + operations.Linear(inner_dim, dim, dtype=dtype, device=device), + ) + + def forward(self, x): + return self.ff(x) + + +class TransformerBlock(nn.Module): + def __init__(self, dim, dim_heads, inner_dim, dtype, device, operations): + super().__init__() + self.pre_norm = LayerNorm(dim, dtype, device) + self.self_attn = Attention(dim, dim_heads, dtype, device, operations) + self.ff_norm = LayerNorm(dim, dtype, device) + self.ff = FeedForward(dim, inner_dim, dtype, device, operations) + + def forward(self, x, rotation_matrix): + x = x + self.self_attn(self.pre_norm(x), rotation_matrix) + return x + self.ff(self.ff_norm(x)) + + +class ContinuousTransformer(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + self.project_in = operations.Linear(2304, 2048, bias=False, dtype=dtype, device=device) + self.project_out = operations.Linear(2048, 128, bias=False, dtype=dtype, device=device) + self.rotary_pos_emb = RotaryEmbedding(32, dtype, device) + self.layers = nn.ModuleList([ + TransformerBlock(2048, 64, 8192, dtype, device, operations) + for _ in range(36) + ]) + + def forward(self, x, timestep_embedding): + x = self.project_in(x) + x = torch.cat((timestep_embedding.unsqueeze(1), x), dim=1) + rotation_matrix = self.rotary_pos_emb.forward_from_seq_len(x.shape[1], x.device, x.dtype) + for layer in self.layers: + x = layer(x, rotation_matrix) + return self.project_out(x[:, 1:]) + + +class DiffusionTransformer(nn.Module): + def __init__(self, dtype, device, operations): + super().__init__() + self.transformer = ContinuousTransformer(dtype, device, operations) + self.timestep_features = FourierFeatures(1, 256, dtype, device) + self.to_timestep_embed = nn.Sequential( + operations.Linear(256, 2048, dtype=dtype, device=device), + nn.SiLU(), + operations.Linear(2048, 2048, dtype=dtype, device=device), + ) + self.preprocess_conv = operations.Conv1d(2304, 2304, 1, bias=False, dtype=dtype, device=device) + self.postprocess_conv = operations.Conv1d(128, 128, 1, bias=False, dtype=dtype, device=device) + + def forward(self, x, timestep, condition): + full = torch.cat((x, torch.zeros_like(x), condition), dim=1) + full = self.preprocess_conv(full) + full + timestep_features = self.timestep_features(timestep[:, None]).to(dtype=x.dtype) + timestep_embedding = self.to_timestep_embed(timestep_features) + out = self.transformer(full.transpose(1, 2), timestep_embedding).transpose(1, 2) + return self.postprocess_conv(out) + out + + +class MiniMaxMusic3DiT(nn.Module): + def __init__(self, dtype=None, device=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.latent_conditioners = nn.Sequential( + operations.Conv1d(4096, 2048, kernel_size=3, padding=1, dtype=dtype, device=device) + ) + self.diffusion_transformer = DiffusionTransformer(dtype, device, operations) + self.cond_layer_logits = nn.Parameter(torch.empty(8, dtype=dtype, device=device)) + self.cond_layer_scale = nn.Parameter(torch.empty(1, dtype=dtype, device=device)) + + def aligned_condition(self, hidden): + frames = hidden.shape[1] + hidden = hidden.transpose(1, 2).reshape(hidden.shape[0], 8, 4096, frames) + weights = torch.softmax(comfy.ops.cast_to_input(self.cond_layer_logits, hidden), dim=0) + hidden = torch.einsum("blht,l->bht", hidden, weights) + hidden = comfy.ops.cast_to_input(self.cond_layer_scale, hidden) * hidden + condition = self.latent_conditioners(hidden) + return torch.nn.functional.interpolate(condition, size=latent_length(frames), mode="nearest") + + def forward(self, x, timestep, context, conditioning_scale, **kwargs): + condition = self.aligned_condition(context) + condition = condition * conditioning_scale[:, :1, :1] + if condition.shape[-1] < x.shape[-1]: + condition = torch.nn.functional.pad(condition, (0, x.shape[-1] - condition.shape[-1])) + else: + condition = condition[..., :x.shape[-1]] + window = latent_length(MAX_CONDITION_FRAMES) + if x.shape[-1] <= window: + return -self.diffusion_transformer(x, timestep, condition) + + output = torch.zeros_like(x) + count = torch.zeros((1, 1, x.shape[-1]), device=x.device, dtype=x.dtype) + hop = latent_length(CONDITION_HOP_FRAMES) + start = 0 + while start < x.shape[-1]: + end = min(start + window, x.shape[-1]) + output[..., start:end] -= self.diffusion_transformer(x[..., start:end], timestep, condition[..., start:end]) + count[..., start:end] += 1 + if end == x.shape[-1]: + break + start += hop + return output / count diff --git a/comfy/ldm/minimax_music/prompt.py b/comfy/ldm/minimax_music/prompt.py new file mode 100644 index 000000000..5f197ee12 --- /dev/null +++ b/comfy/ldm/minimax_music/prompt.py @@ -0,0 +1,70 @@ +import re + + +SPECIAL_TOKEN_IDS = { + "<|im_start|>": 151644, + "<|im_end|>": 151645, + "<|audio_cfg|>": 151654, + "<|audio_start|>": 151669, + "<|audio_end|>": 151670, + "<|caption_start|>": 151671, + "<|caption_end|>": 151672, + "<|lyrics_start|>": 151673, + "<|lyrics_end|>": 151674, +} +AUDIO_CODE_OFFSET = 151675 + +_SPECIAL_TAG_RE = re.compile(r"<\|([^|]*)\|>") +_LYRIC_TAG_RE = re.compile(r"\s*(\[[^\]]+\])\s*") + + +def _remove_markdown_format(text): + lines = [] + for raw_line in text.splitlines(): + line = re.sub(r"^\s{0,3}#{1,6}\s+", "", raw_line) + line = re.sub(r"^\s*[*+-]\s+", "", line) + while "**" in line: + updated = re.sub(r"\*\*([^*]+)\*\*", r"\1", line) + if updated == line: + break + line = updated + line = re.sub(r"(?<|caption_start|>" + f"{clean_caption(caption)}" + "<|caption_end|><|lyrics_start|>" + f"{normalize_lyrics(lyrics)}" + "<|lyrics_end|><|im_end|><|audio_start|>" + ) + + +def validate_tokenizer(tokenizer): + for token, expected in SPECIAL_TOKEN_IDS.items(): + token_id = tokenizer.convert_tokens_to_ids(token) + if token_id != expected: + raise ValueError(f"MiniMax Music3 tokenizer mismatch for {token}: expected {expected}, got {token_id}") diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 2c549e095..b22d03d77 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -10,6 +10,8 @@ from typing import Optional, Any, Callable, Union import logging import functools +import comfy_kitchen + from .diffusionmodules.util import AlphaBlender, timestep_embedding from .sub_quadratic_attention import efficient_dot_product_attention @@ -49,6 +51,8 @@ except ImportError: logging.error(f"\n\nTo use the `--use-flash-attention` feature, the `flash-attn` package must be installed first.\ncommand:\n\t{sys.executable} -m pip install flash-attn") exit(-1) +COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE = comfy_kitchen.int8_attention_is_available() + REGISTERED_ATTENTION_FUNCTIONS = {} def register_attention_function(name: str, func: Callable): # avoid replacing existing functions @@ -145,9 +149,34 @@ def Normalize(in_channels, dtype=None, device=None): return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, dtype=dtype, device=device) +class AttentionTensorContainer: + """Single-owner tensor input consumed by an optimized attention backend.""" + + __slots__ = ("tensor",) + + def __init__(self, tensor: torch.Tensor): + self.tensor: torch.Tensor | None = tensor + + def peek(self) -> torch.Tensor: + if self.tensor is None: + raise RuntimeError("attention tensor container has already been consumed") + return self.tensor + + def take(self) -> torch.Tensor: + tensor = self.peek() + self.tensor = None + return tensor + + def wrap_attn(func): @functools.wraps(func) def wrapper(*args, **kwargs): + containers = None + if len(args) >= 3 and isinstance(args[0], AttentionTensorContainer): + if not isinstance(args[1], AttentionTensorContainer) or not isinstance(args[2], AttentionTensorContainer): + raise TypeError("q, k, and v must all be attention tensor containers") + containers = args[:3] + remove_attn_wrapper_key = False try: if "_inside_attn_wrapper" not in kwargs: @@ -156,11 +185,22 @@ def wrap_attn(func): kwargs["_inside_attn_wrapper"] = True if transformer_options is not None: if "optimized_attention_override" in transformer_options: - return transformer_options["optimized_attention_override"](func, *args, **kwargs) + optimized_attention_override = transformer_options["optimized_attention_override"] + if containers is not None: + if hasattr(optimized_attention_override, "container_function"): + return optimized_attention_override.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] + return optimized_attention_override(func, *args, **kwargs) + + if containers is not None: + if wrapper.container_function is not None: + return wrapper.container_function(*args, **kwargs) + args = tuple(container.take() for container in containers) + args[3:] return func(*args, **kwargs) finally: if remove_attn_wrapper_key: del kwargs["_inside_attn_wrapper"] + wrapper.container_function = None return wrapper @wrap_attn @@ -545,6 +585,63 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha ).transpose(1, 2).reshape(-1, q.shape[2], heads * dim_head) return out +def _comfy_kitchen_int8_inputs(q, k, v, heads, mask, skip_reshape, enable_gqa): + dim_head = q.shape[-1] if skip_reshape else q.shape[-1] // heads + b = q.shape[0] + if not skip_reshape: + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa, expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) + + if mask is not None: + if mask.ndim == 2: + mask = mask.unsqueeze(0) + if mask.ndim == 3: + mask = mask.unsqueeze(1) + + return q, k, v, mask, b, dim_head + + +@wrap_attn +def attention_comfy_kitchen_int8(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + out = comfy_kitchen.int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +def _attention_comfy_kitchen_int8_containers(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): + q = q.take() + k = k.take() + v = v.take() + q, k, v, mask, b, dim_head = _comfy_kitchen_int8_inputs( + q, k, v, heads, mask, skip_reshape, kwargs.get("enable_gqa", False) + ) + quantized = comfy_kitchen.prequantize_int8_attention( + q, + k, + v, + scale=kwargs.get("scale", None), + attn_mask=mask, + ) + del q, k, v + out = comfy_kitchen.int8_attention_from_prequantized(quantized) + if not skip_output_reshape: + out = out.transpose(1, 2).reshape(b, -1, heads * dim_head) + return out + + +attention_comfy_kitchen_int8.container_function = _attention_comfy_kitchen_int8_containers + + @wrap_attn def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape=False, skip_output_reshape=False, **kwargs): if kwargs.get("low_precision_attention", True) is False or (mask is not None and not SAGE_ATTENTION_SUPPORTS_MASK): @@ -775,10 +872,20 @@ else: logging.info("Using sub quadratic optimization for attention, if you have memory or speed issues try using: --use-split-cross-attention") optimized_attention = attention_sub_quad +if model_management.comfy_kitchen_attention_enabled(): + if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + logging.info("Using Comfy Kitchen attention") + optimized_attention = attention_comfy_kitchen_int8 + else: + logging.error("Comfy Kitchen attention is unavailable. Install a Comfy Kitchen build with attention support to use --use-ck-attention.") + exit(-1) + optimized_attention_masked = optimized_attention # register core-supported attention functions +if COMFY_KITCHEN_INT8_ATTENTION_IS_AVAILABLE: + register_attention_function("comfy_kitchen_int8", attention_comfy_kitchen_int8) if SAGE_ATTENTION_IS_AVAILABLE: register_attention_function("sage", attention_sage) if SAGE_ATTENTION3_IS_AVAILABLE: diff --git a/comfy/model_base.py b/comfy/model_base.py index 469d301ea..79f711e92 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -22,6 +22,7 @@ import torch import logging import comfy.ldm.lightricks.av_model import comfy.ldm.minimax.model +import comfy.ldm.minimax_music.dit import comfy.nested_tensor import comfy.ldm.lightricks.symmetric_patchifier import comfy.context_windows @@ -1153,6 +1154,10 @@ class LTXV(BaseModel): if guide_attention_entries is not None: out['guide_attention_entries'] = comfy.conds.CONDConstant(guide_attention_entries) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, **kwargs): @@ -1213,6 +1218,10 @@ class LTXAV(BaseModel): if ref_audio is not None: out['ref_audio'] = comfy.conds.CONDConstant(ref_audio) + generated_keyframes = kwargs.get("generated_keyframes", None) + if generated_keyframes is not None: + out['generated_keyframes'] = comfy.conds.CONDConstant(generated_keyframes) + return out def process_timestep(self, timestep, x, denoise_mask=None, audio_denoise_mask=None, **kwargs): @@ -2156,13 +2165,13 @@ class MiniMaxH3(BaseModel): keyframes = kwargs.get("minimax_keyframes", None) if keyframes is not None: payload["keyframes"] = keyframes - payload["frame_count"] = kwargs.get("minimax_frame_count", None) - payload["cond_video_latents"] = [kf["latent"] for kf in keyframes] + payload["cond_video_latents"] = [kf["latent"] for kf in keyframes if kf.get("latent") is not None] + payload["cond_audio_latents"] = [kf["audio_latent"] for kf in keyframes if kf.get("audio_latent") is not None] refs = kwargs.get("minimax_refs", None) if refs is not None: payload["refs"] = refs - payload["cond_video_latents"] = [r["latent"] for r in refs if "latent" in r] - payload["cond_audio_latents"] = [r["audio_latent"] for r in refs if r.get("audio_latent") is not None] + payload["cond_video_latents"] = payload.get("cond_video_latents", []) + [r["latent"] for r in refs if "latent" in r] + payload["cond_audio_latents"] = payload.get("cond_audio_latents", []) + [r["audio_latent"] for r in refs if r.get("audio_latent") is not None] if kwargs.get("minimax_visual_cond_noise_aug", None) is not None: payload["visual_cond_noise_aug"] = kwargs["minimax_visual_cond_noise_aug"] if kwargs.get("minimax_audio_cond_noise_aug", None) is not None: @@ -2170,16 +2179,80 @@ class MiniMaxH3(BaseModel): payload["seed"] = kwargs.get("seed", 0) # same value process_latent_in/out used, so the model never undoes a scale that was not applied payload["audio_scale"] = self.audio_scale() + + denoise_mask = kwargs.get("denoise_mask", None) + if denoise_mask is not None: + out.update(self._denoise_mask_conds(denoise_mask, latent_shapes)) + if cross_attn is not None and latent_shapes is not None and len(latent_shapes) > 1: # packed layout built once per sampling run, h/w rounded up to the DiT's 2x2 patch vs = latent_shapes[0] payload["layout"] = comfy.ldm.minimax.model.PackedLayout( cross_attn.shape[1], vs[2], (vs[3] + 1) // 2 * 2, (vs[4] + 1) // 2 * 2, latent_shapes[1][-1], keyframes=payload.get("keyframes"), - refs=payload.get("refs"), frame_count=payload.get("frame_count")) + refs=payload.get("refs")) out['minimax_payload'] = comfy.conds.CONDConstant(payload) return out + def _pool_masks_to_token_grid(self, masks): + # pool the per-pixel masks to the label grid with amax: video per 2x2 DiT patch, audio per latent frame + video_mask = masks[0] + h, w = video_mask.shape[-2:] + ph, pw = self.diffusion_model.patch_size[1:] + lead = video_mask.shape[:-2] + video_mask = torch.nn.functional.pad(video_mask.reshape((-1,) + video_mask.shape[-3:]), (0, -w % pw, 0, -h % ph), mode="replicate") + video_mask = video_mask.reshape(lead + video_mask.shape[-2:]) + video_mask = video_mask.reshape(video_mask.shape[:-2] + (video_mask.shape[-2] // ph, ph, video_mask.shape[-1] // pw, pw)).amax(dim=(-3, -1)) + pooled = [video_mask.repeat_interleave(ph, dim=-2).repeat_interleave(pw, dim=-1)[..., :h, :w]] + if len(masks) > 1: + audio_mask = masks[1].amax(dim=1, keepdim=True) + pooled.append(audio_mask.expand_as(masks[1]).contiguous()) + return pooled + + def _token_grid_masks(self, denoise_mask, latent_shapes): + masks = utils.unpack_latents(denoise_mask, latent_shapes) + return [torch.ceil(mask * 256.0) / 256.0 for mask in self._pool_masks_to_token_grid(masks)] + + def _denoise_mask_values(self, denoise_mask, latent_shapes): + if latent_shapes is None or len(latent_shapes) < 2: + return {} + masks = self._token_grid_masks(denoise_mask, latent_shapes) + out = {} + if torch.amin(masks[0]).item() < 1.0 - 1e-3: + out['denoise_mask'] = masks[0][:1, :1].clone() + if torch.amin(masks[1]).item() < 1.0 - 1e-3: + out['audio_denoise_mask'] = masks[1][:1].amax(dim=1, keepdim=True) + return out + + def _denoise_mask_conds(self, denoise_mask, latent_shapes): + return {name: comfy.conds.CONDRegular(value) for name, value in self._denoise_mask_values(denoise_mask, latent_shapes).items()} + + def scale_latent_inpaint(self, sigma, noise, latent_image, x=None, denoise_mask=None, **kwargs): + # preserved regions run at the cond timestep, inject them at cond strength + shapes = self.latent_shapes + if shapes is None or len(shapes) < 2: + return super().scale_latent_inpaint(sigma=sigma, noise=noise, latent_image=latent_image, **kwargs) + cleans = utils.unpack_latents(latent_image, shapes) + noises = utils.unpack_latents(noise, shapes) + aug = comfy.ldm.minimax.model.VISUAL_COND_TIMESTEP # H3's video timestep is 0.999 by default + cleans[0] = aug * cleans[0] + (1.0 - aug) * noises[0] + scale = self.audio_scale() + if scale != 1.0: + # the sampler carries audio as (sigma_v / sigma_a) * x_audio and latent_image + # holds audio_scale * x_audio, so rescale for the model to see it clean + model_sampling = self.model_sampling + sigma_v = sigma.clamp(min=1e-6) + sigma_a = comfy.ldm.minimax.model.time_shift_sigma(sigma_v, model_sampling.shift, model_sampling.audio_shift) + factor = (sigma_v / sigma_a) / scale + cleans[1] = cleans[1] * factor.view(factor.shape[:1] + (1,) * (cleans[1].ndim - 1)).to(cleans[1].dtype) + injected = utils.pack_latents(cleans)[0] + if x is None or denoise_mask is None: + return injected + token_grid_mask = utils.pack_latents(self._token_grid_masks(denoise_mask, shapes))[0] + x_blend_weight = (token_grid_mask - denoise_mask) / (1.0 - denoise_mask).clamp(min=1e-6) + x_blend_weight = torch.where(denoise_mask < 1.0, x_blend_weight.clamp(0.0, 1.0), torch.zeros_like(x_blend_weight)) + return injected + x_blend_weight.to(injected.dtype) * (x - injected) + class TripoSplat(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.triposplat.model.LatentSeqMMFlowModel) @@ -2329,6 +2402,18 @@ class ACEStep15(BaseModel): out['refer_audio'] = comfy.conds.CONDRegular(refer_audio) return out +class MiniMaxMusic3(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.minimax_music.dit.MiniMaxMusic3DiT) + + def process_timestep(self, timestep, **kwargs): + return 1.0 - timestep + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + out["conditioning_scale"] = comfy.conds.CONDRegular(kwargs["conditioning_scale"]) + return out + class Omnigen2(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.omnigen.omnigen2.OmniGen2Transformer2DModel) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 103680fd1..e4bf30b78 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -44,6 +44,13 @@ def calculate_transformer_depth(prefix, state_dict_keys, state_dict): def detect_unet_config(state_dict, key_prefix, metadata=None): state_dict_keys = list(state_dict.keys()) + if ( + '{}cond_layer_logits'.format(key_prefix) in state_dict_keys + and '{}latent_conditioners.0.weight'.format(key_prefix) in state_dict_keys + and '{}diffusion_transformer.transformer.layers.0.self_attn.to_qkv.weight'.format(key_prefix) in state_dict_keys + ): + return {"audio_model": "minimax_music3"} + if '{}joint_blocks.0.context_block.attn.qkv.weight'.format(key_prefix) in state_dict_keys: #mmdit model unet_config = {} unet_config["in_channels"] = state_dict['{}x_embedder.proj.weight'.format(key_prefix)].shape[1] @@ -397,6 +404,7 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["cross_attention_dim"] = shape[1] if metadata is not None and "config" in metadata: dit_config.update(json.loads(metadata["config"]).get("transformer", {})) + dit_config["use_keyframes_abs_pos_embedding"] = '{}keyframes_abs_pos_embedding'.format(key_prefix) in state_dict_keys return dit_config if '{}genre_embedder.weight'.format(key_prefix) in state_dict_keys: #ACE-Step model @@ -829,11 +837,10 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["use_adaln_lora"] = True dit_config["adaln_lora_dim"] = 256 + dit_config["num_blocks"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') if dit_config["model_channels"] == 2048: - dit_config["num_blocks"] = 28 dit_config["num_heads"] = 16 elif dit_config["model_channels"] == 5120: - dit_config["num_blocks"] = 36 dit_config["num_heads"] = 40 if dit_config["in_channels"] == 16: diff --git a/comfy/model_management.py b/comfy/model_management.py index 9f8e7f07b..ff963eb8e 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -490,28 +490,36 @@ try: except: rocm_version = (6, -1) - def aotriton_supported(gpu_arch): - path = torch.__path__[0] - path = os.path.join(os.path.join(path, "lib"), "aotriton.images") - gfx = set(map(lambda a: a[4:], filter(lambda a: a.startswith("amd-gfx"), os.listdir(path)))) - if gpu_arch in gfx: - return True - if "{}x".format(gpu_arch[:-1]) in gfx: - return True - if "{}xx".format(gpu_arch[:-2]) in gfx: - return True - return False + def aotriton_supported(): + """Whether pytorch reports flash attention as usable on this gpu. + + can_use_flash_attention() evaluates runtime eligibility for the given + parameters; on a ROCm build that includes checking the gpu arch against the + kernel images AOTriton was compiled for. Querying it avoids assuming where + those images live inside the torch install. The probe tensor is shaped and + typed to pass the unrelated SDPA checks, so False means no hardware support + rather than a rejected shape. + """ + try: + if not torch.backends.cuda.is_flash_attention_available(): # not built with flash attention + return False + q = torch.empty((1, 1, 8, 64), dtype=torch.float16, device=get_torch_device()) + params = torch.backends.cuda.SDPAParams(q, q, q, None, 0.0, False, False) + return torch.backends.cuda.can_use_flash_attention(params, False) + except (AttributeError, RuntimeError, TypeError) as e: + logging.warning("Could not query aotriton support: {}".format(e)) + return False logging.info("AMD arch: {}".format(arch)) logging.info("ROCm version: {}".format(rocm_version)) if args.use_split_cross_attention == False and args.use_quad_cross_attention == False: - if aotriton_supported(arch): # AMD efficient attention implementation depends on aotriton. + if aotriton_supported(): # AMD efficient attention implementation depends on aotriton. if torch_version_numeric >= (2, 7): # works on 2.6 but doesn't actually seem to improve much if any((a in arch) for a in ["gfx90a", "gfx942", "gfx950", "gfx1100", "gfx1101", "gfx1150", "gfx1151"]): # TODO: more arches, TODO: gfx950 ENABLE_PYTORCH_ATTENTION = True if rocm_version >= (7, 0): - if any((a in arch) for a in ["gfx1200", "gfx1201"]): - ENABLE_PYTORCH_ATTENTION = True + if any((a in arch) for a in ["gfx1200", "gfx1201"]): + ENABLE_PYTORCH_ATTENTION = True if torch_version_numeric >= (2, 7) and rocm_version >= (6, 4): if any((a in arch) for a in ["gfx1200", "gfx1201", "gfx950"]): # TODO: more arches, "gfx942" gives error on pytorch nightly 2.10 1013 rocm7.0 SUPPORT_FP8_OPS = True @@ -1360,9 +1368,14 @@ STREAM_CAST_BUFFERS = {} LARGEST_CASTED_WEIGHT = (None, 0) STREAM_AIMDO_CAST_BUFFERS = {} LARGEST_AIMDO_CASTED_WEIGHT = (None, 0) +CROSS_STEP_STATE = weakref.WeakSet() DEFAULT_AIMDO_CAST_BUFFER_RESERVATION_SIZE = 16 * 1024 ** 3 +# NOTE: devs/agents: this is temporary and will be removed in a future comfy. Not supported for custom node use. +def _register_cross_step(module): + CROSS_STEP_STATE.add(module) + def get_cast_buffer(offload_stream, device, size, ref): global LARGEST_CASTED_WEIGHT @@ -1417,6 +1430,10 @@ def reset_cast_buffers(): mmap_obj.bounce() DIRTY_MMAPS.clear() + for module in CROSS_STEP_STATE: + del module._comfy_cross_step_state + CROSS_STEP_STATE.clear() + for loaded_model in current_loaded_models: model = loaded_model.model if model is not None and model.is_dynamic(): @@ -1658,6 +1675,9 @@ def unpin_memory(tensor): def sage_attention_enabled(): return args.use_sage_attention +def comfy_kitchen_attention_enabled(): + return args.use_ck_attention + def flash_attention_enabled(): return args.use_flash_attention diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index ae3f0191d..72942aa04 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -685,6 +685,14 @@ class ModelPatcher: def set_model_attn2_output_patch(self, patch): self.set_model_patch(patch, "attn2_output_patch") + def set_model_optimized_attention(self, optimized_attention): + def optimized_attention_override(_, *args, **kwargs): + return optimized_attention(*args, **kwargs) + + if hasattr(optimized_attention, "container_function") and optimized_attention.container_function is not None: + optimized_attention_override.container_function = optimized_attention.container_function + self.model_options["transformer_options"]["optimized_attention_override"] = optimized_attention_override + def set_model_input_block_patch(self, patch): self.set_model_patch(patch, "input_block_patch") @@ -1879,8 +1887,29 @@ class ModelPatcherDynamic(ModelPatcher): loading = self._load_list(for_dynamic=True, default_device=device_to) loading.sort() + get_units = getattr(self.model, "get_dynamic_vram__units", None) + dynamic_units, last_dynamic_units = get_units() if get_units is not None else ([], []) + dynamic_units = list(dynamic_units) + last_dynamic_units = list(last_dynamic_units) + loading_by_module = {entry[-2]: entry for entry in loading} + loading = [] + for unit in dynamic_units: + unit_modules = unit if isinstance(unit, (list, tuple)) else (unit,) + modules = [module for root in unit_modules for module in root.modules() if module in loading_by_module] + for index, module in enumerate(modules): + loading.append((*loading_by_module.pop(module), unit if index == len(modules) - 1 else None)) + last_loading = [] + for unit in last_dynamic_units: + unit_modules = unit if isinstance(unit, (list, tuple)) else (unit,) + modules = [module for root in unit_modules for module in root.modules() if module in loading_by_module] + for index, module in enumerate(modules): + last_loading.append((*loading_by_module.pop(module), unit if index == len(modules) - 1 else None)) + loading.extend((*entry, None) for entry in loading_by_module.values()) + loading.extend(last_loading) + v_block = None + for x in loading: - *_, module_mem, n, m, params = x + *_, module_mem, n, m, params, end_of_block = x def set_dirty(item, dirty): if dirty or not hasattr(item, "_v_signature"): @@ -1973,6 +2002,13 @@ class ModelPatcherDynamic(ModelPatcher): move_weight_functions(m, device_to) + if hasattr(m, "_v"): + v_block = m._v if v_block is None else (v_block[0], v_block[1], max(v_block[2], m._v[1] + m._v[2] - v_block[1])) + if end_of_block is not None: + unit = end_of_block + (unit[0] if isinstance(unit, (list, tuple)) else unit)._v_block = v_block + v_block = None + for key, buf in self.model.named_buffers(recurse=True): if key not in self.backup_buffers: self.backup_buffers[key] = buf diff --git a/comfy/model_prefetch.py b/comfy/model_prefetch.py index aa6d22d77..bdde5137a 100644 --- a/comfy/model_prefetch.py +++ b/comfy/model_prefetch.py @@ -1,11 +1,19 @@ +import torch +import warnings +import weakref + import comfy_aimdo.model_vbar +from comfy.cli_args import args import comfy.memory_management import comfy.model_management import comfy.ops PREFETCH_QUEUES = [] +GRAPH_MODULES = weakref.WeakSet() +GRAPH_WARMED_MODULES = weakref.WeakSet() +GRAPH_CAPTURE_STREAMS = {} -def cleanup_prefetched_modules(comfy_modules): +def cleanup_prefetched_modules(module, comfy_modules): for s in comfy_modules: prefetch = getattr(s, "_prefetch", None) if prefetch is None: @@ -17,39 +25,86 @@ def cleanup_prefetched_modules(comfy_modules): if prefetch["signature"] is not None: comfy_aimdo.model_vbar.vbar_unpin(s._v) delattr(s, "_prefetch") + if getattr(module, "_v_block_faulted", False): + comfy_aimdo.model_vbar.vbar_unpin(module._v_block) + del module._v_block_faulted + +def _drop_graph(module): + graph = getattr(module, "_comfy_graph", None) + if graph is None: + return + # reset() through the bound method surfaces the allocator's benign + # "uncaptured free of a captured allocation" as catchable Python warnings; + # a plain del frees from the C++ dealloc path and spams stderr instead + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + graph["graph"].reset() + del module._comfy_graph def cleanup_prefetch_queues(): - global PREFETCH_QUEUES + global PREFETCH_QUEUES, GRAPH_CAPTURE_STREAMS for queue in PREFETCH_QUEUES: for entry in queue: if entry is None or not isinstance(entry, tuple): continue _, prefetch_state = entry - comfy_modules = prefetch_state[1] + prefetched_module, comfy_modules = prefetch_state if comfy_modules is not None: - cleanup_prefetched_modules(comfy_modules) + cleanup_prefetched_modules(prefetched_module, comfy_modules) PREFETCH_QUEUES = [] + for module in GRAPH_MODULES: + _drop_graph(module) + GRAPH_MODULES.clear() + GRAPH_WARMED_MODULES.clear() + GRAPH_CAPTURE_STREAMS = {} -def prefetch_queue_pop(queue, device, module): +def prefetch_queue_pop(queue, device, module, dtype=None, core=None, enable_graph=False, generator=None): + enable_graph = enable_graph and not args.disable_cuda_graphs and comfy.model_management.is_device_cuda(device) and getattr(module, "_v_block", None) is not None if queue is None: + if core is not None: + core() return + capture_stream = None + if enable_graph: + capture_stream = GRAPH_CAPTURE_STREAMS.get(device) + if capture_stream is None: + capture_stream = torch.cuda.Stream(device=device) + GRAPH_CAPTURE_STREAMS[device] = capture_stream + + signature = None + graph_hit = False + graph = getattr(module, "_comfy_graph", None) if enable_graph else None + if graph is not None: + signature = comfy_aimdo.model_vbar.vbar_fault(module._v_block) + if signature is not None: + module._v_block_faulted = True + graph_hit = comfy_aimdo.model_vbar.vbar_signature_compare(signature, graph["signature"]) + consumed = queue.pop(0) if consumed is not None: offload_stream, prefetch_state = consumed if offload_stream is not None: offload_stream.wait_stream(comfy.model_management.current_stream(device)) - _, comfy_modules = prefetch_state + prefetched_module, comfy_modules = prefetch_state if comfy_modules is not None: - cleanup_prefetched_modules(comfy_modules) + cleanup_prefetched_modules(prefetched_module, comfy_modules) + if graph_hit: + queue[0] = (None, (module, [])) + graph["graph"].replay() + return + + fully_faulted = False prefetch = queue[0] if prefetch is not None: comfy_modules = [] - for s in prefetch.modules(): - if hasattr(s, "_v"): - comfy_modules.append(s) + prefetch_modules = prefetch if isinstance(prefetch, (list, tuple)) else (prefetch,) + for root in prefetch_modules: + for s in root.modules(): + if hasattr(s, "_v"): + comfy_modules.append(s) registerable_size = 0 for s in comfy_modules: @@ -59,11 +114,42 @@ def prefetch_queue_pop(queue, device, module): if lowvram_fn is not None: registerable_size += lowvram_fn.memory_required() - offload_stream = comfy.ops.cast_modules_with_vbar(comfy_modules, None, device, None, True) + offload_stream, fully_faulted = comfy.ops.cast_modules_with_vbar(comfy_modules, None, device, None, True, return_faulted=True) if not comfy.model_management.args.fast_disk: comfy.model_management.ensure_pin_registerable(registerable_size) comfy.model_management.sync_stream(device, offload_stream) - queue[0] = (offload_stream, (prefetch, comfy_modules)) + if fully_faulted and dtype is not None: + for comfy_module in comfy_modules: + comfy.ops.resolve_cast_module_with_vbar(comfy_module, dtype, device, dtype, None, False, return_weights=False) + queue[0] = (offload_stream, (module, comfy_modules)) + + if core is not None: + if enable_graph and fully_faulted and module in GRAPH_WARMED_MODULES: + if signature is None: + signature = comfy_aimdo.model_vbar.vbar_fault(module._v_block) + if signature is not None: + module._v_block_faulted = True + if signature is not None: + _drop_graph(module) + graph = torch.cuda.CUDAGraph() + if generator is not None: + graph.register_generator_state(generator) + capture_stream.wait_stream(comfy.model_management.current_stream(device)) + with torch.cuda.graph(graph, stream=capture_stream, capture_error_mode="thread_local"): + core() + comfy.model_management.current_stream(device).wait_stream(capture_stream) + graph.replay() + module._comfy_graph = {"graph": graph, "signature": signature} + GRAPH_MODULES.add(module) + return + if capture_stream is None: + core() + else: + capture_stream.wait_stream(comfy.model_management.current_stream(device)) + with torch.cuda.stream(capture_stream): + core() + comfy.model_management.current_stream(device).wait_stream(capture_stream) + GRAPH_WARMED_MODULES.add(module) def make_prefetch_queue(queue, device, transformer_options): if (not transformer_options.get("prefetch_dynamic_vbars", False) diff --git a/comfy/nested_tensor.py b/comfy/nested_tensor.py index 08c7133f8..43835b15f 100644 --- a/comfy/nested_tensor.py +++ b/comfy/nested_tensor.py @@ -83,6 +83,9 @@ class NestedTensor: def layout(self): return self.tensors[0].layout + def __repr__(self): + return f"{type(self).__name__}({self.tensors!r})" + def cat_nested(tensors, *args, **kwargs): cated_tensors = [] diff --git a/comfy/ops.py b/comfy/ops.py index 14599997b..ff64aad59 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -123,10 +123,12 @@ def materialize_meta_param(s, param_keys): # FIXME: add n=1 cache hit fast path -def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blocking): +def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blocking, return_faulted=False): offload_stream = None cast_buffer = None cast_buffer_offset = 0 + if return_faulted: + fully_faulted = all(not getattr(s, param_key + "_function", []) for s in comfy_modules for param_key in ("weight", "bias")) def ensure_offload_stream(module, required_size, check_largest): nonlocal offload_stream @@ -163,6 +165,8 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin for s in comfy_modules: signature = comfy_aimdo.model_vbar.vbar_fault(s._v) resident = comfy_aimdo.model_vbar.vbar_signature_compare(signature, s._v_signature) + if return_faulted and (signature is None or not resident): + fully_faulted = False prefetch = { "signature": signature, "resident": resident, @@ -255,10 +259,12 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin prefetch["needs_cast"] = needs_cast s._prefetch = prefetch + if return_faulted: + return offload_stream, fully_faulted return offload_stream -def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, want_requant): +def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, want_requant, return_weights=True): prefetch = getattr(s, "_prefetch", None) @@ -298,7 +304,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w tensor = tensor.dequantize() return tensor - if orig.dtype != dtype or len(fns) > 0: + if (return_weights and orig.dtype != dtype) or len(fns) > 0: x = to_dequant(x, dtype) if not resident and lowvram_fn is not None: x = to_dequant(x, dtype if compute_dtype is None else compute_dtype) @@ -325,7 +331,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w if prefetch["signature"] is not None: prefetch["resident"] = True - return weight, bias + return (weight, bias) if return_weights else None def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None, offloadable=False, compute_dtype=None, want_requant=False): @@ -452,6 +458,26 @@ def uncast_bias_weight(s, weight, bias, offload_stream): device = bias_a.device os.wait_stream(comfy.model_management.current_stream(device)) +class CastBiasWeightContext: + # When initialized with no arguments or the first is None, the context + # will return the tuple (None, None). + def __init__(self, *args, **kwargs): + self.slf = args[0] if len(args) else None + self.state = (None, None) if self.slf is None else cast_bias_weight(*args, **kwargs) + + def __enter__(self): + result = self.state + if len(result) < 3 or result[2] is None: + # Not offloaded, immediately drop references. + self.state = self.slf = None + return result[:2] + + def __exit__(self, *_args) -> None: + if self.slf is None: + return + slf, state = self.slf, self.state + self.state = self.slf = None + uncast_bias_weight(slf, *state) class CastWeightBiasOp: comfy_cast_weights = False @@ -538,10 +564,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.linear(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.linear(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -555,10 +579,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -572,10 +594,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias) def forward(self, *args, **kwargs): run_every_op() @@ -600,10 +620,8 @@ class disable_weight_init: return super()._conv_forward(input, weight, bias, *args, **kwargs) def forward_comfy_cast_weights(self, input, autopad=None): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = self._conv_forward(input, weight, bias, autopad=autopad) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return self._conv_forward(input, weight, bias, autopad=autopad) def forward(self, *args, **kwargs): run_every_op() @@ -617,10 +635,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -634,12 +650,10 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - running_mean = self.running_mean.to(device=input.device, dtype=weight.dtype) if self.running_mean is not None else None - running_var = self.running_var.to(device=input.device, dtype=weight.dtype) if self.running_var is not None else None - x = torch.nn.functional.batch_norm(input, running_mean, running_var, weight, bias, self.training, self.momentum, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + running_mean = self.running_mean.to(device=input.device, dtype=weight.dtype) if self.running_mean is not None else None + running_var = self.running_var.to(device=input.device, dtype=weight.dtype) if self.running_var is not None else None + return torch.nn.functional.batch_norm(input, running_mean, running_var, weight, bias, self.training, self.momentum, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -653,15 +667,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - if self.weight is not None: - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - else: - weight = None - bias = None - offload_stream = None - x = torch.nn.functional.layer_norm(input, self.normalized_shape, weight, bias, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self if self.weight is not None else None, input, offloadable=True) as (weight, bias): + return torch.nn.functional.layer_norm(input, self.normalized_shape, weight, bias, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -676,15 +683,8 @@ class disable_weight_init: return None def forward_comfy_cast_weights(self, input): - if self.weight is not None: - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - else: - weight = None - bias = None - offload_stream = None - x = torch.nn.functional.rms_norm(input, self.normalized_shape, weight, self.eps) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self if self.weight is not None else None, input, offloadable=True) as (weight, bias): + return torch.nn.functional.rms_norm(input, self.normalized_shape, weight, self.eps) def forward(self, *args, **kwargs): run_every_op() @@ -703,12 +703,10 @@ class disable_weight_init: input, output_size, self.stride, self.padding, self.kernel_size, num_spatial_dims, self.dilation) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.conv_transpose2d( - input, weight, bias, self.stride, self.padding, - output_padding, self.groups, self.dilation) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.conv_transpose2d( + input, weight, bias, self.stride, self.padding, + output_padding, self.groups, self.dilation) def forward(self, *args, **kwargs): run_every_op() @@ -727,12 +725,10 @@ class disable_weight_init: input, output_size, self.stride, self.padding, self.kernel_size, num_spatial_dims, self.dilation) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.conv_transpose1d( - input, weight, bias, self.stride, self.padding, - output_padding, self.groups, self.dilation) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.conv_transpose1d( + input, weight, bias, self.stride, self.padding, + output_padding, self.groups, self.dilation) def forward(self, *args, **kwargs): run_every_op() @@ -795,10 +791,8 @@ class disable_weight_init: output_dtype = out_dtype if self.weight.dtype == torch.float16 or self.weight.dtype == torch.bfloat16: out_dtype = None - weight, bias, offload_stream = cast_bias_weight(self, device=input.device, dtype=out_dtype, offloadable=True) - x = torch.nn.functional.embedding(input, weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse).to(dtype=output_dtype) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, device=input.device, dtype=out_dtype, offloadable=True) as (weight, bias): + return torch.nn.functional.embedding(input, weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse).to(dtype=output_dtype) def forward(self, *args, **kwargs): @@ -874,7 +868,6 @@ def fp8_linear(self, input): if input.ndim != 2: return None lora_compute_dtype=comfy.model_management.lora_compute_dtype(input.device) - w, bias, offload_stream = cast_bias_weight(self, input, dtype=dtype, bias_dtype=input_dtype, offloadable=True, compute_dtype=lora_compute_dtype, want_requant=True) scale_weight = torch.ones((), device=input.device, dtype=torch.float32) scale_input = torch.ones((), device=input.device, dtype=torch.float32) @@ -883,15 +876,16 @@ def fp8_linear(self, input): layout_params_input = TensorCoreFP8Layout.Params(scale=scale_input, orig_dtype=input_dtype, orig_shape=tuple(input_fp8.shape)) quantized_input = QuantizedTensor(input_fp8, "TensorCoreFP8Layout", layout_params_input) - # Wrap weight in QuantizedTensor - this enables unified dispatch - # Call F.linear - __torch_dispatch__ routes to fp8_linear handler in quant_ops.py! - layout_params_weight = TensorCoreFP8Layout.Params(scale=scale_weight, orig_dtype=input_dtype, orig_shape=tuple(w.shape)) - quantized_weight = QuantizedTensor(w, "TensorCoreFP8Layout", layout_params_weight) - o = torch.nn.functional.linear(quantized_input, quantized_weight, bias) + with CastBiasWeightContext(self, input, dtype=dtype, bias_dtype=input_dtype, offloadable=True, compute_dtype=lora_compute_dtype, want_requant=True) as (w, bias): + # Wrap weight in QuantizedTensor - this enables unified dispatch + # Call F.linear - __torch_dispatch__ routes to fp8_linear handler in quant_ops.py! + w_shape = tuple(w.shape) + layout_params_weight = TensorCoreFP8Layout.Params(scale=scale_weight, orig_dtype=input_dtype, orig_shape=w_shape) + quantized_weight = QuantizedTensor(w, "TensorCoreFP8Layout", layout_params_weight) + o = torch.nn.functional.linear(quantized_input, quantized_weight, bias) - uncast_bias_weight(self, w, bias, offload_stream) if tensor_3d: - o = o.reshape((input_shape[0], input_shape[1], w.shape[0])) + o = o.reshape((input_shape[0], input_shape[1], w_shape[0])) return o @@ -911,10 +905,8 @@ class fp8_ops(manual_cast): except Exception as e: logging.info("Exception during fp8 op: {}".format(e)) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = torch.nn.functional.linear(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return torch.nn.functional.linear(input, weight, bias) CUBLAS_IS_AVAILABLE = False try: @@ -930,10 +922,8 @@ if CUBLAS_IS_AVAILABLE: return None def forward_comfy_cast_weights(self, input): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - x = cublas_half_matmul(input, weight, bias, self._epilogue_str, self.has_bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): + return cublas_half_matmul(input, weight, bias, self._epilogue_str, self.has_bias) def forward(self, *args, **kwargs): run_every_op() @@ -1344,29 +1334,28 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec want_requant=False, weight_only_quant=False, ): - if weight_only_quant: - weight, bias, offload_stream = cast_bias_weight( - self, - input=None, - dtype=self.weight.dtype, - device=input.device, - bias_dtype=input.dtype, - offloadable=True, - compute_dtype=compute_dtype, - want_requant=True, - ) - weight = weight.to(dtype=input.dtype) - else: - weight, bias, offload_stream = cast_bias_weight( + if not weight_only_quant: + with CastBiasWeightContext( self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=want_requant, - ) - x = self._forward(input, weight, bias) - uncast_bias_weight(self, weight, bias, offload_stream) - return x + ) as (weight, bias): + return self._forward(input, weight, bias) + + with CastBiasWeightContext( + self, + input=None, + dtype=self.weight.dtype, + device=input.device, + bias_dtype=input.dtype, + offloadable=True, + compute_dtype=compute_dtype, + want_requant=True, + ) as (weight, bias): + weight = weight.to(dtype=input.dtype) + return self._forward(input, weight, bias) def forward(self, input, *args, **kwargs): run_every_op() @@ -1391,25 +1380,20 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec # Training path: quantized forward with compute_dtype backward via autograd function if (input.requires_grad and _use_quantized and quantize_input): - - weight, bias, offload_stream = cast_bias_weight( + with CastBiasWeightContext( self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=True - ) + ) as (weight, bias): + scale = getattr(self, 'input_scale', None) + if scale is not None: + scale = comfy.model_management.cast_to_device(scale, input.device, None) - scale = getattr(self, 'input_scale', None) - if scale is not None: - scale = comfy.model_management.cast_to_device(scale, input.device, None) - - output = QuantLinearFunc.apply( - input, weight, bias, self.layout_type, scale, compute_dtype - ) - - uncast_bias_weight(self, weight, bias, offload_stream) - return output + return QuantLinearFunc.apply( + input, weight, bias, self.layout_type, scale, compute_dtype + ) # Inference path (unchanged) if _use_quantized and quantize_input: @@ -1520,13 +1504,11 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec """Cast the whole bank once; expert_linear inside reuses the cast. Not re-entrant — do not nest calls on the same instance. """ - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - self._resident_bank = (weight, bias) - try: - yield self - finally: - self._resident_bank = None - uncast_bias_weight(self, weight, bias, offload_stream) + with CastBiasWeightContext(self, input, offloadable=True) as self._resident_bank: + try: + yield self + finally: + self._resident_bank = None def expert_linear(self, input: torch.Tensor, i: int) -> torch.Tensor: """Linear against expert i's weight (with optional bias).""" @@ -1534,11 +1516,8 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec if resident is not None: weight, bias = resident return self._expert_linear_impl(input, weight, bias, i) - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True) - try: + with CastBiasWeightContext(self, input, offloadable=True) as (weight, bias): return self._expert_linear_impl(input, weight, bias, i) - finally: - uncast_bias_weight(self, weight, bias, offload_stream) def _expert_linear_impl(self, input, weight, bias, i): if isinstance(weight, QuantizedTensor): @@ -1641,28 +1620,26 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec # Optimized path: lookup in fp8/int8, dequantize only the selected rows. if isinstance(weight, QuantizedTensor) and len(self.weight_function) == 0: - qdata, _, offload_stream = cast_bias_weight(self, device=input.device, dtype=weight.dtype, offloadable=True) - if isinstance(qdata, QuantizedTensor): - params = qdata._params - scale = params.scale - qdata = qdata._qdata - else: - params = weight._params - scale = None + with CastBiasWeightContext(self, device=input.device, dtype=weight.dtype, offloadable=True) as (qdata, _bias): + if isinstance(qdata, QuantizedTensor): + params = qdata._params + scale = params.scale + qdata = qdata._qdata + else: + params = weight._params + scale = None - # int8: per-row scale possible ConvRot, so let the layout do the gather - if self.quant_format == "int8_tensorwise": - x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input) - uncast_bias_weight(self, qdata, None, offload_stream) - return x if out_dtype is None else x.to(dtype=out_dtype) + # int8: per-row scale possible ConvRot, so let the layout do the gather + if self.quant_format == "int8_tensorwise": + x = get_layout_class(self.layout_type).dequantize_embedding(qdata, params, input) + return x if out_dtype is None else x.to(dtype=out_dtype) - x = torch.nn.functional.embedding( - input, qdata, self.padding_idx, self.max_norm, - self.norm_type, self.scale_grad_by_freq, self.sparse) - uncast_bias_weight(self, qdata, None, offload_stream) + x = torch.nn.functional.embedding( + input, qdata, self.padding_idx, self.max_norm, + self.norm_type, self.scale_grad_by_freq, self.sparse) target_dtype = out_dtype if out_dtype is not None else weight._params.orig_dtype x = x.to(dtype=target_dtype) - if scale is not None and scale != 1.0: + if scale is not None: x = x * scale.to(dtype=target_dtype) return x diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index 6d9112dbb..18fd2d613 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -40,7 +40,7 @@ try: cuda_version = tuple(map(int, str(torch.version.cuda).split('.'))) if cuda_version < (13,): ck.registry.disable("cuda") - logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.") + logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.\nWARNING WARNING WARNING\nIf you are on nvidia 20 series and above it is required that you update your pytorch to cu130 or higher.\n") # 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 diff --git a/comfy/sample.py b/comfy/sample.py index 2be0cae5f..4ebad9666 100644 --- a/comfy/sample.py +++ b/comfy/sample.py @@ -37,10 +37,15 @@ def prepare_noise(latent_image, seed, noise_inds=None): return noises +def prepare_empty_noise(latent_image): + if latent_image.is_nested: + return comfy.nested_tensor.NestedTensor([torch.zeros_like(t, device="cpu") for t in latent_image.unbind()]) + return torch.zeros_like(latent_image, device="cpu") + def fix_empty_latent_channels(model, latent_image, downscale_ratio_spacial=None, downscale_ratio_temporal=None): if latent_image.is_nested: return latent_image - latent_format = model.get_model_object("latent_format") #Resize the empty latent image so it has the right number of channels + latent_format = model.get_model_object("latent_format") is_empty = torch.count_nonzero(latent_image) == 0 if is_empty: if latent_format.latent_channels != latent_image.shape[1]: @@ -59,6 +64,9 @@ def fix_empty_latent_channels(model, latent_image, downscale_ratio_spacial=None, new_t = max(1, round(latent_image.shape[2] * ratio)) latent_image = comfy.utils.repeat_to_batch_size(latent_image, new_t, dim=2) + if is_empty: + latent_image = latent_format.fix_empty_latent(latent_image) + return latent_image def prepare_sampling(model, noise_shape, positive, negative, noise_mask): diff --git a/comfy/samplers.py b/comfy/samplers.py index 1d6a4e104..94307c1a7 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -636,7 +636,7 @@ class KSamplerX0Inpaint: if "denoise_mask_function" in model_options: denoise_mask = model_options["denoise_mask_function"](sigma, denoise_mask, extra_options={"model": self.inner_model, "sigmas": self.sigmas}) latent_mask = 1. - denoise_mask - x = x * denoise_mask + self.inner_model.inner_model.scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image) * latent_mask + x = x * denoise_mask + self.inner_model.inner_model.scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image, denoise_mask=denoise_mask) * latent_mask out = self.inner_model(x, sigma, model_options=model_options, seed=seed) if denoise_mask is not None: out = out * denoise_mask + self.latent_image * latent_mask diff --git a/comfy/sd.py b/comfy/sd.py index 9ccd561bc..06679c6fb 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -11,6 +11,7 @@ from .ldm.cascade.stage_c_coder import StageC_coder from .ldm.audio.autoencoder import AudioOobleckVAE import comfy.ldm.genmo.vae.model import comfy.ldm.lightricks.vae.causal_video_autoencoder +import comfy.ldm.lightricks.vae.na_diffusion_decoder import comfy.ldm.lightricks.vae.audio_vae import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae @@ -24,6 +25,7 @@ import comfy.ldm.cogvideo.vae import comfy.ldm.hunyuan_video.vae import comfy.ldm.mmaudio.vae.autoencoder import comfy.ldm.audio.vae_sa3 +import comfy.ldm.minimax_music.dav import comfy.pixel_space_convert import comfy.weight_adapter import yaml @@ -31,6 +33,7 @@ import math import os import comfy.utils +import comfy.ops from . import clip_vision from . import gligen @@ -73,6 +76,7 @@ import comfy.text_encoders.longcat_image import comfy.text_encoders.qwen35 import comfy.text_encoders.qwen3vl import comfy.text_encoders.minimax +import comfy.text_encoders.minimax_music import comfy.ldm.minimax.vae import comfy.ldm.minimax.audio_vae import comfy.text_encoders.boogu @@ -514,7 +518,22 @@ class VAE: self.audio_sample_rate = 44100 if config is None: - if "decoder.mid.block_1.mix_factor" in sd: + if "dec_in_proj.weight" in sd and "decoder.model.0.weight_g" in sd: # MiniMax Music3 DAV + self.first_stage_model = comfy.ldm.minimax_music.dav.MiniMaxMusic3DAV(operations=comfy.ops.disable_weight_init) + self.latent_channels = 128 + self.output_channels = 2 + self.upscale_ratio = 512 + self.downscale_ratio = 512 + self.latent_dim = 1 + self.process_output = lambda audio: audio + self.process_input = lambda audio: audio + self.working_dtypes = [torch.float32] + self.disable_offload = True + self.memory_used_decode = lambda shape, dtype: (shape[-1] * 512 * 1400 + 800_000_000) * model_management.dtype_size(dtype) + def _no_encode(*args, **kwargs): + raise RuntimeError("MiniMax Music3 DAV cannot encode audio") + self.memory_used_encode = _no_encode + elif "decoder.mid.block_1.mix_factor" in sd: encoder_config = {'double_z': True, 'z_channels': 4, 'resolution': 256, 'in_channels': 3, 'out_ch': 3, 'ch': 128, 'ch_mult': [1, 2, 4, 4], 'num_res_blocks': 2, 'attn_resolutions': [], 'dropout': 0.0} decoder_config = encoder_config.copy() decoder_config["video_kernel_size"] = [3, 1, 1] @@ -583,6 +602,22 @@ class VAE: self.working_dtypes = [torch.bfloat16, torch.float32] self.memory_used_encode = lambda shape, dtype: (400 * shape[2] * shape[3]) * model_management.dtype_size(dtype) self.memory_used_decode = lambda shape, dtype: (1000 * shape[2] * shape[3] * 16 * 16) * model_management.dtype_size(dtype) + elif "decoder.conv_in_x_t.weight" in sd: # lightricks LTX 2.4 diffusion VAE decoder + vae_config = None + if metadata is not None and "config" in metadata: + vae_config = json.loads(metadata["config"]).get("vae", None) + self.first_stage_model = comfy.ldm.lightricks.vae.na_diffusion_decoder.CausalDiffusionVAE(config=vae_config) + self.latent_channels = sd["decoder.conv_in.weight"].shape[1] + self.latent_dim = 3 + self.disable_offload = True + self.crop_input = False # generic crop would narrow the frame axis by the 32x spatial ratio + self.memory_used_decode = lambda shape, dtype: (1700 * shape[2] * shape[3] * shape[4] * (8 * 8 * 8)) * model_management.dtype_size(dtype) + self.memory_used_encode = lambda shape, dtype: (80 * max(shape[2], 7) * shape[3] * shape[4]) * model_management.dtype_size(dtype) + self.upscale_ratio = (lambda a: max(0, a * 8 - 7), 32, 32) + self.upscale_index_formula = (8, 32, 32) + self.downscale_ratio = (lambda a: max(0, math.floor((a + 7) / 8)), 32, 32) + self.downscale_index_formula = (8, 32, 32) + self.working_dtypes = [torch.bfloat16, torch.float32] 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} @@ -872,7 +907,14 @@ class VAE: self.upscale_index_formula = (4, 16, 16) self.downscale_ratio = (lambda a: max(0, math.floor((a + 3) / 4)), 16, 16) self.downscale_index_formula = (4, 16, 16) - if self.latent_channels in [48, 128]: # Wan 2.2 and LTX2 + if self.latent_channels == 24 and sd["decoder.22.bias"].shape[0] == 12: # MiniMax H3 + self.first_stage_model = comfy.taesd.taehv.TAEHV(latent_channels=self.latent_channels, latent_format=None) + self.process_input = self.process_output = lambda image: image + self.upscale_ratio = (lambda a: max(1, (a - 2) // 5 * 17 + 5), 16, 16) + self.downscale_ratio = (lambda a: max(1, (a - 1) // 17 * 5 + 2) if a > 1 else 1, 16, 16) + self.memory_used_encode = lambda shape, dtype: (400 * ((shape[-3] + 16) // 17) * shape[-2] * shape[-1] * model_management.dtype_size(dtype)) + self.memory_used_decode = lambda shape, dtype: ((260 * 16 * 16 + shape[1] * shape[-3]) * shape[-2] * shape[-1] * model_management.dtype_size(dtype)) + elif self.latent_channels in [48, 128]: # Wan 2.2 and LTX2 self.first_stage_model = comfy.taesd.taehv.TAEHV(latent_channels=self.latent_channels, latent_format=None) # taehv doesn't need scaling self.process_input = self.process_output = lambda image: image self.process_output = lambda image: image @@ -955,13 +997,21 @@ class VAE: self.working_dtypes = [torch.float16, torch.float32] # the model tiles internally (256px spatial, 17-frame temporal chunks) self.handles_tiling = True + # decode finalizes straight to [0, 1] while streaming chunks out + self.process_output = lambda image: image + # one decoded temporal chunk (with overlap) is all that ever sits in VRAM + chunk_frames = (self.first_stage_model.tokens_chunk_size + self.first_stage_model.token_overlap) * self.first_stage_model.vae_ratio_t + def estimate_encode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames == 1 else 1_300_000_000 elements_per_pixel = 7 if frames == 1 else 9.5 + # only one clip of the input video is ever resident on the GPU + frames = min(frames, self.first_stage_model.clip_length) return (elements_per_pixel * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 def estimate_decode_memory(frames, height, width, dtype): fixed = 110_000_000 if frames <= 22 else 270_000_000 + frames = min(frames, chunk_frames + 2) return (9.5 * frames * height * width + fixed) * model_management.dtype_size(dtype) * 1.03 self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], shape[3], shape[4], dtype) @@ -1198,6 +1248,7 @@ class VAE: do_tile = True if do_tile: + pixel_samples = None comfy.model_management.soft_empty_cache() dims = samples_in.ndim - 2 if dims == 1 or self.extra_1d_channel is not None: @@ -1213,16 +1264,48 @@ class VAE: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 if self.handles_tiling: + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, tile, tile, None), self.vae_dtype) + model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) 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)) + # Reserve as much as an untiled decode could use (capped by what the device can provide), then size the tiles to fill that reservation: + # shrink the temporal tile until one tile fits, then grow the spatial tile while it still fits. + budget = min(memory_used, int(model_management.get_total_memory(self.device) * 0.8)) + model_management.load_models_gpu([self.patcher], memory_required=budget, force_full_load=self.disable_offload) + tile_t = samples_in.shape[2] + est = lambda tt, txy: self.memory_used_decode(self._tile_bounded_shape(samples_in.shape, txy, txy, tt), self.vae_dtype) + while tile_t > 2 and est(tile_t, tile) > budget: + tile_t = -(-tile_t // 2) + while tile * 2 <= max(samples_in.shape[3], samples_in.shape[4]) and est(tile_t, tile * 2) <= budget: + tile *= 2 + overlap = tile // 4 + pixel_samples = self.decode_tiled_3d(samples_in, tile_t=tile_t, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples + def _tile_bounded_shape(self, shape, tile_x, tile_y, tile_t): + """Clamp a latent shape to one tile for memory estimates: peak memory of a tiled decode is per-tile. Only caller-provided tile dims are clamped.""" + s = list(shape) + if len(s) == 5: + if tile_t is not None: + s[2] = min(s[2], tile_t) + if tile_y is not None: + s[3] = min(s[3], tile_y) + if tile_x is not None: + s[4] = min(s[4], tile_x) + elif len(s) == 4 and self.extra_1d_channel is None: + if tile_y is not None: + s[2] = min(s[2], tile_y) + if tile_x is not None: + s[3] = min(s[3], tile_x) + elif tile_x is not None: + s[-1] = min(s[-1], tile_x) + return tuple(s) + def decode_tiled(self, samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): self.throw_exception_if_invalid() - memory_used = self.memory_used_decode(samples.shape, self.vae_dtype) #TODO: calculate mem required for tile + memory_used = self.memory_used_decode(self._tile_bounded_shape(samples.shape, tile_x, tile_y, tile_t), self.vae_dtype) model_management.load_models_gpu([self.patcher], memory_required=memory_used, force_full_load=self.disable_offload) dims = samples.ndim - 2 args = {} @@ -1634,7 +1717,16 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.params = {} if len(clip_data) == 1: te_model = detect_te_model(clip_data[0]) - if te_model == TEModel.CLIP_G: + if clip_type == CLIPType.MINIMAX and "model.audio_decoder.projection.weight" in clip_data[0]: + tokenizer_data["tokenizer_json"] = clip_data[0].pop("tokenizer_json", None) + quant = comfy.utils.detect_layer_quantization(clip_data[0], "") + if quant is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = quant + clip_target.params["projection_config"] = comfy.text_encoders.minimax_music.detect_merged_config(clip_data[0]) + clip_target.clip = comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel + clip_target.tokenizer = comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer + elif te_model == TEModel.CLIP_G: if clip_type == CLIPType.STABLE_CASCADE: clip_target.clip = sdxl_clip.StableCascadeClipModel clip_target.tokenizer = sdxl_clip.StableCascadeTokenizer @@ -1693,12 +1785,21 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B): - variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, - TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, - TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, - TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] - clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) - clip_target.tokenizer = variant.tokenizer + if te_model == TEModel.GEMMA_4_12B and "text_embedding_projection.video_aggregate_embed.weight" in clip_data[0]: + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(comfy.text_encoders.gemma4.Gemma4_12B), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(comfy.text_encoders.gemma4.Gemma4_12B.tokenizer) + else: + variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model] + clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant) + clip_target.tokenizer = variant.tokenizer tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None) elif te_model == TEModel.GEMMA_2_2B: if clip_type == CLIPType.PIXELDIT: @@ -1866,9 +1967,30 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.clip = comfy.text_encoders.kandinsky5.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.kandinsky5.Kandinsky5TokenizerImage elif clip_type == CLIPType.LTXV: - clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) - clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer - tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) + te_models = [detect_te_model(sd) for sd in clip_data] + gemma4_models = { + TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B, + TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B, + TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B, + TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B, + } + gemma4_type = next((model for model in te_models if model in gemma4_models), None) + if gemma4_type is None: + clip_target.clip = comfy.text_encoders.lt.ltxav_te(**llama_detect(clip_data), **comfy.text_encoders.lt.sd_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.lt.LTXAVGemmaTokenizer + gemma_sd = clip_data[te_models.index(TEModel.GEMMA_3_12B)] if TEModel.GEMMA_3_12B in te_models else clip_data[0] + tokenizer_data["spiece_model"] = gemma_sd.get("spiece_model", None) + else: + variant = gemma4_models[gemma4_type] + clip_target.clip = comfy.text_encoders.lt.ltxav_te( + **llama_detect(clip_data), + **comfy.text_encoders.lt.sd_detect(clip_data), + text_encoder_model=comfy.text_encoders.gemma4.gemma4_text_encoder_model(variant), + text_encoder_key="gemma4", + ) + clip_target.tokenizer = comfy.text_encoders.lt.ltxav_gemma4_tokenizer(variant.tokenizer) + gemma_sd = clip_data[te_models.index(gemma4_type)] + tokenizer_data["tokenizer_json"] = gemma_sd.get("tokenizer_json", None) elif clip_type == CLIPType.NEWBIE: clip_target.clip = comfy.text_encoders.newbie.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.newbie.NewBieTokenizer diff --git a/comfy/supported_models.py b/comfy/supported_models.py index b9952db55..33c378435 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -16,6 +16,7 @@ import comfy.text_encoders.genmo import comfy.text_encoders.lt import comfy.text_encoders.hunyuan_video import comfy.text_encoders.minimax +import comfy.text_encoders.minimax_music import comfy.text_encoders.cosmos import comfy.text_encoders.lumina2 import comfy.text_encoders.wan @@ -2200,6 +2201,28 @@ class ACEStep15(supported_models_base.BASE): return supported_models_base.ClipTarget(comfy.text_encoders.ace15.ACE15Tokenizer, comfy.text_encoders.ace15.te(**detect)) +class MiniMaxMusic3(supported_models_base.BASE): + unet_config = { + "audio_model": "minimax_music3", + } + + latent_format = comfy.latent_formats.MiniMaxMusic3 + memory_usage_factor = 2.0 + supported_inference_dtypes = [torch.float16, torch.bfloat16, torch.float32] + sampling_settings = {"multiplier": 1.0} + + def get_model(self, state_dict, prefix="", device=None): + return model_base.MiniMaxMusic3(self, device=device) + + def model_type(self, state_dict, prefix=""): + return model_base.ModelType.FLOW + + def clip_target(self, state_dict={}): + detect = comfy.text_encoders.minimax_music.detect_merged_config(state_dict, self.text_encoder_key_prefix[0]) + target = supported_models_base.ClipTarget(comfy.text_encoders.minimax_music.MiniMaxMusic3Tokenizer, comfy.text_encoders.minimax_music.MiniMaxMusic3TEModel) + target.params["projection_config"] = detect + return target + class LongCatImage(supported_models_base.BASE): unet_config = { @@ -2494,6 +2517,7 @@ models = [ ChromaRadiance, ACEStep, ACEStep15, + MiniMaxMusic3, Omnigen2, Boogu, MageFlow, diff --git a/comfy/taesd/taehv.py b/comfy/taesd/taehv.py index 696013200..ffa9f89d1 100644 --- a/comfy/taesd/taehv.py +++ b/comfy/taesd/taehv.py @@ -131,10 +131,11 @@ class TAEHV(nn.Module): self.latent_channels = latent_channels self.parallel = parallel self.latent_format = latent_format + self.is_h3 = self.latent_channels == 24 self.show_progress_bar = show_progress_bar self.process_in = latent_format().process_in if latent_format is not None else (lambda x: x) self.process_out = latent_format().process_out if latent_format is not None else (lambda x: x) - if self.latent_channels in [48, 32]: # Wan 2.2 and HunyuanVideo1.5 + if self.latent_channels in [48, 32, 24]: # Wan 2.2, HunyuanVideo1.5 and MiniMax H3 self.patch_size = 2 elif self.latent_channels == 128: # LTX2 self.patch_size, self.latent_channels, encoder_time_downscale, decoder_time_upscale = 4, 128, (True, True, True), (True, True, True) @@ -176,6 +177,21 @@ class TAEHV(nn.Module): def encode(self, x, **kwargs): x = x.movedim(2, 1) # [B, C, T, H, W] -> [B, T, C, H, W] + if self.is_h3: + single_frame = x.shape[1] == 1 + batch = x.shape[0] + x = torch.cat([x, x[:, -1:].expand(-1, -x.shape[1] % 17, -1, -1, -1)], dim=1) + x = F.pad(x.reshape(batch, -1, 17, *x.shape[2:]), (0, 0, 0, 0, 0, 0, 3, 0)) + if self.parallel: + x = apply_model_with_memblocks(self.encoder, x.flatten(0, 1), True, self.show_progress_bar, + patch_size=self.patch_size) + x = x.reshape(batch, -1, *x.shape[2:]) + else: + x = torch.cat([apply_model_with_memblocks(self.encoder, chunk, False, False, + patch_size=self.patch_size) + for chunk in tqdm(x.unbind(1), disable=not self.show_progress_bar)], dim=1) + x = x[:, :1] if single_frame else x[:, :-3] + return self.process_out(x.movedim(2, 1)) if x.shape[1] % self.t_downscale != 0: # pad at end to multiple of t_downscale n_pad = self.t_downscale - x.shape[1] % self.t_downscale @@ -189,7 +205,16 @@ class TAEHV(nn.Module): x = x.unsqueeze(0) if x.ndim == 4 else x # [T, C, H, W] -> [1, T, C, H, W] x = x.movedim(1, 2) if x.shape[1] != self.latent_channels else x # [B, T, C, H, W] or [B, C, T, H, W] x = self.process_in(x).movedim(2, 1) # [B, C, T, H, W] -> [B, T, C, H, W] + if self.is_h3: + single_frame = x.shape[1] == 1 x = apply_model_with_memblocks(self.decoder, x, self.parallel, self.show_progress_bar, output_device=comfy.model_management.intermediate_device(), patch_size=self.patch_size, decode=True) + if self.is_h3: + x.clamp_(0, 1) + if not single_frame: + chunk_frames = 5 * self.t_upscale + x = F.pad(x, (0, 0, 0, 0, 0, 0, 0, -x.shape[1] % chunk_frames)) + x = x.unflatten(1, (-1, chunk_frames))[:, :, self.frames_to_trim:].flatten(1, 2) + return x[:, :-3 * self.t_upscale].movedim(2, 1) return x[:, self.frames_to_trim:].movedim(2, 1) diff --git a/comfy/text_encoders/ace15.py b/comfy/text_encoders/ace15.py index 853f021ae..3ad519314 100644 --- a/comfy/text_encoders/ace15.py +++ b/comfy/text_encoders/ace15.py @@ -4,9 +4,29 @@ from comfy import sd1_clip import torch import math import yaml +import comfy.ops import comfy.utils +def _audio_logits(model, x, audio_start, audio_end, eos_token=None): + input = x[:, -1:] + module = model.embed_tokens + + offload_stream = None + if module.comfy_cast_weights: + weight, _, offload_stream = comfy.ops.cast_bias_weight(module, input, offloadable=True) + else: + weight = module.weight.to(x) + + logits = torch.nn.functional.linear(input, weight[audio_start:audio_end], None)[:, -1] + eos_logits = None + if eos_token is not None: + eos_logits = torch.nn.functional.linear(input, weight[eos_token:eos_token + 1], None)[:, -1] + + comfy.ops.uncast_bias_weight(module, weight, None, offload_stream) + return logits, eos_logits + + def sample_manual_loop_no_classes( model, ids=None, @@ -34,48 +54,43 @@ def sample_manual_loop_no_classes( execution_dtype = torch.float32 embeds, attention_mask, num_tokens, embeds_info = model.process_tokens(ids, device) + embeds = embeds.to(execution_dtype) embeds_batch = embeds.shape[0] - output_audio_codes = [] - past_key_values = [] + output_audio_codes = torch.empty((max_new_tokens,), device=device, dtype=torch.long) + generated_tokens = 0 generator = torch.Generator(device=device) generator.manual_seed(seed) - model_config = model.transformer.model.config - past_kv_shape = [embeds_batch, model_config.num_key_value_heads, embeds.shape[1] + min_tokens, model_config.head_dim] - - for x in range(model_config.num_hidden_layers): - past_key_values.append((torch.empty(past_kv_shape, device=device, dtype=execution_dtype), torch.empty(past_kv_shape, device=device, dtype=execution_dtype), 0)) + past_key_values = model.transformer.model.init_kv_cache(embeds_batch, embeds.shape[1] + max_new_tokens, device, execution_dtype) + fixed_kv = isinstance(past_key_values[0], comfy.text_encoders.llama.FixedKV) progress_bar = comfy.utils.ProgressBar(max_new_tokens) + sampling_logits = None for step in comfy.utils.model_trange(max_new_tokens, desc="LM sampling"): - outputs = model.transformer(None, attention_mask, embeds=embeds.to(execution_dtype), num_tokens=num_tokens, intermediate_output=None, dtype=execution_dtype, embeds_info=embeds_info, past_key_values=past_key_values) - next_token_logits = model.transformer.logits(outputs[0])[:, -1] + outputs = model.transformer(None, attention_mask, embeds=embeds, num_tokens=num_tokens, intermediate_output=None, dtype=execution_dtype, embeds_info=embeds_info, past_key_values=past_key_values) past_key_values = outputs[2] - if cfg_scale != 1.0: - cond_logits = next_token_logits[0:1] - uncond_logits = next_token_logits[1:2] - cfg_logits = uncond_logits + cfg_scale * (cond_logits - uncond_logits) - else: - cfg_logits = next_token_logits[0:1] - use_eos_score = eos_token_id is not None and eos_token_id < audio_start_id and min_tokens < step - if use_eos_score: - eos_score = cfg_logits[:, eos_token_id].clone() + audio_logits, eos_logits = _audio_logits(model.transformer.model, outputs[0], audio_start_id, audio_end_id, eos_token_id if use_eos_score else None) + if cfg_scale != 1.0: + cfg_logits = audio_logits[1:2] + cfg_scale * (audio_logits[0:1] - audio_logits[1:2]) + if use_eos_score: + cond_eos = eos_logits[0:1, 0] + uncond_eos = eos_logits[1:2, 0] + eos_score = uncond_eos + cfg_scale * (cond_eos - uncond_eos) + else: + cfg_logits = audio_logits[0:1] + if use_eos_score: + eos_score = eos_logits[0:1, 0] remove_logit_value = torch.finfo(cfg_logits.dtype).min - # Only generate audio tokens - cfg_logits[:, :audio_start_id] = remove_logit_value - cfg_logits[:, audio_end_id:] = remove_logit_value - if use_eos_score: - cfg_logits[:, eos_token_id] = eos_score + cfg_logits = torch.cat((eos_score.unsqueeze(1), cfg_logits), dim=1) if top_k is not None and top_k > 0: - top_k_vals, _ = torch.topk(cfg_logits, top_k) - min_val = top_k_vals[..., -1, None] - cfg_logits[cfg_logits < min_val] = remove_logit_value + top_k_values = torch.topk(cfg_logits, min(top_k, cfg_logits.shape[-1])).values + cfg_logits[cfg_logits < top_k_values[..., -1, None]] = remove_logit_value if min_p is not None and min_p > 0: probs = torch.softmax(cfg_logits, dim=-1) @@ -89,28 +104,40 @@ def sample_manual_loop_no_classes( sorted_indices_to_remove = cumulative_probs > top_p sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] = 0 - indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove) + indices_to_remove = torch.zeros_like(cfg_logits, dtype=torch.bool) + indices_to_remove.scatter_(1, sorted_indices, sorted_indices_to_remove) cfg_logits[indices_to_remove] = remove_logit_value if temperature > 0: cfg_logits = cfg_logits / temperature - next_token = torch.multinomial(torch.softmax(cfg_logits, dim=-1), num_samples=1, generator=generator).squeeze(1) + if sampling_logits is None: + sampling_logits = cfg_logits.new_empty((cfg_logits.shape[0], model.transformer.model.vocab_size)) + sampling_logits.fill_(remove_logit_value) + if use_eos_score: + sampling_logits[:, eos_token_id] = cfg_logits[:, 0] + cfg_logits = cfg_logits[:, 1:] + sampling_logits[:, audio_start_id:audio_end_id] = cfg_logits + next_token = torch.multinomial(torch.softmax(sampling_logits, dim=-1), num_samples=1, generator=generator).squeeze(1) else: next_token = torch.argmax(cfg_logits, dim=-1) + if use_eos_score: + next_token = torch.where(next_token == 0, eos_token_id, next_token + audio_start_id - 1) + else: + next_token += audio_start_id - token = next_token.item() - - if token == eos_token_id: + if eos_token_id is not None and next_token.item() == eos_token_id: break - embed, _, _, _ = model.process_tokens([[token]], device) - embeds = embed.repeat(embeds_batch, 1, 1) - attention_mask = torch.cat([attention_mask, torch.ones((embeds_batch, 1), device=device, dtype=attention_mask.dtype)], dim=1) + input_ids = next_token.repeat(embeds_batch).unsqueeze(1) + embeds = model.transformer.get_input_embeddings()(input_ids, out_dtype=execution_dtype) + if not fixed_kv: + attention_mask = torch.cat([attention_mask, torch.ones((embeds_batch, 1), device=device, dtype=attention_mask.dtype)], dim=1) - output_audio_codes.append(token - audio_start_id) + output_audio_codes[generated_tokens] = next_token[0] - audio_start_id + generated_tokens += 1 progress_bar.update_absolute(step) - return output_audio_codes + return output_audio_codes[:generated_tokens].tolist() def generate_audio_codes(model, positive, negative, min_tokens=1, max_tokens=1024, seed=0, cfg_scale=2.0, temperature=0.85, top_p=0.9, top_k=0, min_p=0.000): @@ -286,7 +313,10 @@ class ACE15TEModel(torch.nn.Module): self.qwen3_06b = Qwen3_06BModel(device=device, dtype=dtype, model_options=model_options) if model is not None: setattr(self, self.lm_model, model(device=device, dtype=dtype_llama, model_options=model_options)) - + ar_model = getattr(self, self.lm_model) + ar_model.transformer.model.fixed_kv = True + ar_model.transformer.model.prefetch_dynamic_vbars = True + ar_model.transformer.model.graph_dynamic_vbar_blocks = True self.dtypes = set([dtype, dtype_llama]) def encode_token_weights(self, token_weight_pairs): @@ -319,6 +349,12 @@ class ACE15TEModel(torch.nn.Module): if lm_model is not None: lm_model.reset_clip_options() + def get_dynamic_vram__units(self): + if self.lm_model is None: + return ([], []) + model = getattr(self, self.lm_model) + return model.transformer.model.get_dynamic_vram__units() + def load_sd(self, sd): if "model.layers.0.post_attention_layernorm.weight" in sd: shape = sd["model.layers.0.post_attention_layernorm.weight"].shape diff --git a/comfy/text_encoders/bpe_tokenizer.py b/comfy/text_encoders/bpe_tokenizer.py new file mode 100644 index 000000000..e49e36ca0 --- /dev/null +++ b/comfy/text_encoders/bpe_tokenizer.py @@ -0,0 +1,333 @@ +""" +Pure-Python byte-level BPE tokenizer. +Supports loading from HuggingFace tokenizer.json (LLaMA-style) +and from Mistral tekken JSON blobs. +No dependency on the `transformers`, `tokenizers`, or `regex` packages. +""" +import base64 +import json +import os +import re +import unicodedata + + +# This is also the default pattern used by the previous MistralConverter path. +_LLAMA_PATTERN = r"""(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+""" +_CONTRACTIONS = ("'re", "'ve", "'ll", "'s", "'t", "'m", "'d") + + +def _is_letter(c): + return unicodedata.category(c)[0] == "L" + + +def _is_number(c): + return unicodedata.category(c)[0] == "N" + + +def _is_whitespace(c): + return c in " \t\n\r\v\f\x85\u2028\u2029" or unicodedata.category(c) == "Zs" + + +def _split_llama(text): + pieces = [] + i = 0 + while i < len(text): + contraction = None + if text[i] == "'": + for suffix in _CONTRACTIONS: + if text[i:i + len(suffix)].casefold() == suffix: + contraction = text[i:i + len(suffix)] + break + if contraction is not None: + pieces.append(contraction) + i += len(contraction) + continue + + j = i + if text[j] not in "\r\n" and not _is_letter(text[j]) and not _is_number(text[j]): + j += 1 + if j < len(text) and _is_letter(text[j]): + j += 1 + while j < len(text) and _is_letter(text[j]): + j += 1 + pieces.append(text[i:j]) + i = j + continue + + if _is_number(text[i]): + j = i + 1 + while j < len(text) and j - i < 3 and _is_number(text[j]): + j += 1 + pieces.append(text[i:j]) + i = j + continue + + j = i + if text[j] == " ": + j += 1 + punct_start = j + while j < len(text) and not _is_whitespace(text[j]) and not _is_letter(text[j]) and not _is_number(text[j]): + j += 1 + if j > punct_start: + while j < len(text) and text[j] in "\r\n": + j += 1 + pieces.append(text[i:j]) + i = j + continue + + if _is_whitespace(text[i]): + j = i + 1 + while j < len(text) and _is_whitespace(text[j]): + j += 1 + last_newline = max(text.rfind("\r", i, j), text.rfind("\n", i, j)) + if last_newline >= i: + j = last_newline + 1 + elif j < len(text) and j - i > 1: + j -= 1 + pieces.append(text[i:j]) + i = j + continue + + pieces.append(text[i]) + i += 1 + return pieces + + +def _make_split_pattern(pattern_str): + if pattern_str != _LLAMA_PATTERN: + raise ValueError(f"Unsupported tokenizer split pattern: {pattern_str}") + return _split_llama + + +def _bytes_to_unicode(): + bs = (list(range(ord("!"), ord("~") + 1)) + + list(range(ord("¡"), ord("¬") + 1)) + + list(range(ord("®"), ord("ÿ") + 1))) + cs = bs[:] + n = 0 + for b in range(2**8): + if b not in bs: + bs.append(b) + cs.append(2**8 + n) + n += 1 + cs = [chr(n) for n in cs] + return dict(zip(bs, cs)) + + +class BPETokenizer: + """Byte-level BPE tokenizer with optional BOS prepending.""" + + def __init__(self, vocab, merges_by_pair, special_token_ids, pattern_str, + byte_encoder, byte_decoder, bos_id=None): + self._vocab = vocab # str -> int + self._inv_vocab = {v: k for k, v in vocab.items()} + self._merges = merges_by_pair # (str, str) -> priority int + self._special_token_ids = special_token_ids # str -> int + self._special_ids = set(special_token_ids.values()) + self._byte_encoder = byte_encoder + self._byte_decoder = byte_decoder + self._bos_id = bos_id + + self._split = _make_split_pattern(pattern_str) + sorted_specials = sorted(special_token_ids.keys(), key=len, reverse=True) + if sorted_specials: + self._special_split = re.compile( + '(' + '|'.join(re.escape(s) for s in sorted_specials) + ')' + ) + else: + self._special_split = None + + def _bpe_encode_piece(self, chars): + if len(chars) <= 1: + return chars + while True: + min_rank = float('inf') + best_pair = None + for i in range(len(chars) - 1): + r = self._merges.get((chars[i], chars[i + 1]), float('inf')) + if r < min_rank: + min_rank = r + best_pair = (chars[i], chars[i + 1]) + if best_pair is None: + break + merged = best_pair[0] + best_pair[1] + new_chars = [] + i = 0 + while i < len(chars): + if i < len(chars) - 1 and chars[i] == best_pair[0] and chars[i + 1] == best_pair[1]: + new_chars.append(merged) + i += 2 + else: + new_chars.append(chars[i]) + i += 1 + chars = new_chars + if len(chars) == 1: + break + return chars + + def _encode_raw(self, text): + ids = [] + parts = self._special_split.split(text) if self._special_split else [text] + for part in parts: + if not part: + continue + if part in self._special_token_ids: + ids.append(self._special_token_ids[part]) + else: + for piece in self._split(part): + byte_chars = [self._byte_encoder[b] for b in piece.encode('utf-8')] + for tok in self._bpe_encode_piece(byte_chars): + ids.append(self._vocab[tok]) + return ids + + def __call__(self, text): + ids = self._encode_raw(text) + if self._bos_id is not None: + ids = [self._bos_id] + ids + return {"input_ids": ids} + + def get_vocab(self): + return dict(self._vocab) + + def decode(self, token_ids, skip_special_tokens=True): + buf = bytearray() + for tid in token_ids: + s = self._inv_vocab.get(tid, '') + if tid in self._special_ids: + if not skip_special_tokens: + buf.extend(s.encode('utf-8')) + else: + for c in s: + buf.append(self._byte_decoder[c]) + return buf.decode('utf-8', errors='replace') + + +def _extract_pattern(pretok): + if pretok.get('type') == 'Sequence': + for sub in pretok.get('pretokenizers', []): + if sub.get('type') == 'Split': + pat = sub.get('pattern', {}) + if 'Regex' in pat: + return pat['Regex'] + elif pretok.get('type') == 'Split': + pat = pretok.get('pattern', {}) + if 'Regex' in pat: + return pat['Regex'] + return None + + +def _extract_bos_id(post_processor, special_token_ids): + if post_processor.get('type') == 'TemplateProcessing': + single = post_processor.get('single', []) + if single and 'SpecialToken' in single[0]: + bos_str = single[0]['SpecialToken']['id'] + return special_token_ids.get(bos_str) + return None + + +def from_tokenizer_json(path): + """Load a BPETokenizer from a directory containing tokenizer.json.""" + tok_file = os.path.join(path, 'tokenizer.json') + with open(tok_file, encoding='utf-8') as f: + data = json.load(f) + + vocab = dict(data['model']['vocab']) # str -> int + + merges_by_pair = {} + for i, merge_str in enumerate(data['model'].get('merges', [])): + a, b = merge_str.split(' ', 1) + if (a, b) not in merges_by_pair: + merges_by_pair[(a, b)] = i + + special_token_ids = {} + for tok in data.get('added_tokens', []): + special_token_ids[tok['content']] = tok['id'] + vocab[tok['content']] = tok['id'] # include in vocab for inv_vocab decode + + pattern = _extract_pattern(data.get('pre_tokenizer', {})) + if pattern is None: + raise ValueError(f"Could not extract regex pattern from {tok_file}") + + bos_id = _extract_bos_id(data.get('post_processor', {}), special_token_ids) + + byte_encoder = _bytes_to_unicode() + byte_decoder = {v: k for k, v in byte_encoder.items()} + + return BPETokenizer(vocab, merges_by_pair, special_token_ids, pattern, + byte_encoder, byte_decoder, bos_id=bos_id) + + +def from_tekken_json(data): + """Build a BPETokenizer from a Mistral tekken JSON blob (bytes or str).""" + mistral_vocab = json.loads(data) + config = mistral_vocab["config"] + + byte_encoder = _bytes_to_unicode() + byte_decoder = {v: k for k, v in byte_encoder.items()} + + def tbts(b): + return "".join(byte_encoder[ord(c)] for c in b.decode("latin-1")) + + special_token_offset = config["default_num_special_tokens"] + max_vocab = config["default_vocab_size"] - special_token_offset + + raw_vocab = {} + for w in mistral_vocab["vocab"]: + r = w["rank"] + if r >= max_vocab: + continue + raw_vocab[base64.b64decode(w["token_bytes"])] = r + special_token_offset + + special_tokens_dict = {} + for w in mistral_vocab["special_tokens"]: + if "token_bytes" in w: + special_tokens_dict[base64.b64decode(w["token_bytes"])] = w["rank"] + else: + special_tokens_dict[w["token_str"]] = w["rank"] + + all_special = list(special_tokens_dict.keys()) + combined = dict(special_tokens_dict) + combined.update(raw_vocab) + + bpe_vocab = {} + merge_triples = [] + for token, rank in combined.items(): + if token not in all_special: + bpe_vocab[tbts(token)] = rank + if len(token) == 1: + continue + local = [] + for i in range(1, len(token)): + pl, pr = token[:i], token[i:] + if pl in combined and pr in combined and (pl + pr) in combined: + local.append((pl, pr, rank)) + local.sort(key=lambda x: (combined[x[0]], combined[x[1]])) + merge_triples.extend(local) + else: + tok_str = token.decode("utf-8", errors="replace") if isinstance(token, bytes) else token + bpe_vocab[tok_str] = rank + + merge_triples.sort(key=lambda v: v[2]) + + merges_by_pair = {} + for i, (pl, pr, _) in enumerate(merge_triples): + pair = (tbts(pl), tbts(pr)) + if pair not in merges_by_pair: + merges_by_pair[pair] = i + + special_str_ids = {} + for tok in all_special: + tok_str = tok.decode("utf-8", errors="replace") if isinstance(tok, bytes) else tok + if tok_str in bpe_vocab: + special_str_ids[tok_str] = bpe_vocab[tok_str] + + return BPETokenizer(bpe_vocab, merges_by_pair, special_str_ids, _LLAMA_PATTERN, + byte_encoder, byte_decoder, bos_id=None) + + +class LlamaTokenizerFast: + """Drop-in replacement for transformers.LlamaTokenizerFast (read-only use).""" + + @staticmethod + def from_pretrained(path, **kwargs): + return from_tokenizer_json(path) diff --git a/comfy/text_encoders/flux.py b/comfy/text_encoders/flux.py index d5eb91dcb..fbdb1d13a 100644 --- a/comfy/text_encoders/flux.py +++ b/comfy/text_encoders/flux.py @@ -3,11 +3,10 @@ import comfy.text_encoders.t5 import comfy.text_encoders.sd3_clip import comfy.text_encoders.llama import comfy.model_management -from transformers import T5TokenizerFast, LlamaTokenizerFast, Qwen2Tokenizer +from transformers import T5TokenizerFast, Qwen2Tokenizer +from .bpe_tokenizer import from_tekken_json import torch import os -import json -import base64 class T5XXLTokenizer(sd1_clip.SDTokenizer): def __init__(self, embedding_directory=None, tokenizer_data={}): @@ -75,45 +74,13 @@ def flux_clip(dtype_t5=None, t5_quantization_metadata=None): def load_mistral_tokenizer(data): if torch.is_tensor(data): data = data.numpy().tobytes() + return {"tokenizer_object": from_tekken_json(data)} - try: - from transformers.integrations.mistral import MistralConverter - except ModuleNotFoundError: - from transformers.models.pixtral.convert_pixtral_weights_to_hf import MistralConverter - - mistral_vocab = json.loads(data) - - special_tokens = {} - vocab = {} - - max_vocab = mistral_vocab["config"]["default_vocab_size"] - max_vocab -= len(mistral_vocab["special_tokens"]) - - for w in mistral_vocab["vocab"]: - r = w["rank"] - if r >= max_vocab: - continue - - vocab[base64.b64decode(w["token_bytes"])] = r - - for w in mistral_vocab["special_tokens"]: - if "token_bytes" in w: - special_tokens[base64.b64decode(w["token_bytes"])] = w["rank"] - else: - special_tokens[w["token_str"]] = w["rank"] - - all_special = [] - for v in special_tokens: - all_special.append(v) - - special_tokens.update(vocab) - vocab = special_tokens - return {"tokenizer_object": MistralConverter(vocab=vocab, additional_special_tokens=all_special).converted(), "legacy": False} class MistralTokenizerClass: @staticmethod - def from_pretrained(path, **kwargs): - return LlamaTokenizerFast(**kwargs) + def from_pretrained(path, tokenizer_object=None, **kwargs): + return tokenizer_object class Mistral3Tokenizer(sd1_clip.SDTokenizer): def __init__(self, embedding_directory=None, embedding_size=5120, embedding_key='mistral3_24b', tokenizer_data={}): diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 5163c1676..0d8f0fcc7 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -6,13 +6,16 @@ import numpy as np from tokenizers import Tokenizer from dataclasses import dataclass import math +import re from comfy import sd1_clip import comfy.model_management +import comfy.model_prefetch import comfy.ops +import comfy.quant_ops from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.rmsnorm import rms_norm -from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding +from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, FixedKV, _make_scaled_embedding # Intentional minor divergences from transformers -reference implementation: @@ -109,7 +112,28 @@ class Gemma4_12B_Config(Gemma4Config): suppress_tokens = [258883, 258882] -# unfused RoPE as addcmul_ RoPE diverges from reference code +class RingKV(FixedKV): + # sliding-window ring: writes wrap at capacity, validity saturates + def prepare(self, num_tokens): + capacity = self.key.shape[2] + self.position.fill_(self.index % capacity) + self.seqlen.fill_(min(self.index + num_tokens, capacity)) + + +def _fixed_kv_decode_mask(mask, cache, min_val): + capacity = cache.key.shape[2] + valid = min(cache.index + 1, capacity) + output = mask.new_full((*mask.shape[:-1], capacity), min_val) + if isinstance(cache, RingKV): + positions = torch.arange(cache.index + 1 - valid, cache.index + 1, device=mask.device) % capacity + output.index_copy_(-1, positions, mask[..., -valid:]) + else: + output[..., :valid] = mask[..., :valid] + return output + + +# unfused RoPE as addcmul_ RoPE diverges from reference code (vision only; text +# layers use the kitchen split-half kernel, bitwise-equal to this with bf16 freqs) def _apply_rotary_pos_emb(x, freqs_cis): cos, sin = freqs_cis[0], freqs_cis[1] half = x.shape[-1] // 2 @@ -140,6 +164,23 @@ class Gemma4Attention(nn.Module): if config.k_norm == "gemma3": self.k_norm = RMSNorm(head_dim, eps=config.rms_norm_eps, device=device, dtype=dtype) + def _decode_attention(self, xq, cache, bias): + if bias is None: + # eager decode: slice the cache to the valid length (python-side index, + # no mask needed; a full ring is order-invariant under softmax) + n = min(cache.index + 1, cache.key.shape[2]) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + attention = optimized_attention_for_device(xq.device, mask=False, small_input=True) + return attention(xq, cache.key[:, :, :n], cache.value[:, :, :n], self.num_heads, skip_reshape=True, scale=1.0, **gqa_kwargs) + # graph capture: fixed-length masked attention over the full capacity, explicit + # math (SDPA leaves its fast path on broadcast-bias + GQA and costs ~0.5ms/layer) + batch_size = xq.shape[0] + groups = self.num_heads // self.num_kv_heads + q = xq.reshape(batch_size, self.num_kv_heads, groups, self.head_dim) + scores = q @ cache.key.transpose(-1, -2) + bias + probs = torch.softmax(scores.float(), dim=-1).to(xq.dtype) + return (probs @ cache.value).reshape(batch_size, 1, self.inner_size) + def forward( self, hidden_states: torch.Tensor, @@ -156,10 +197,16 @@ class Gemma4Attention(nn.Module): if self.q_norm is not None: xq = self.q_norm(xq) + if isinstance(shared_kv, FixedKV): + # decode on a KV-shared layer: attend the source layer's fixed cache + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) + output = self._decode_attention(xq, shared_kv, attention_mask) + return self.o_proj(output), None, None + if shared_kv is not None: xk, xv = shared_kv # Apply RoPE to Q only (K already has RoPE from source layer) - xq = _apply_rotary_pos_emb(xq, freqs_cis) + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) present_key_value = None shareable_kv = None else: @@ -173,11 +220,40 @@ class Gemma4Attention(nn.Module): xv = rms_norm(xv) xk = xk.transpose(1, 2) xv = xv.transpose(1, 2) - xq = _apply_rotary_pos_emb(xq, freqs_cis) - xk = _apply_rotary_pos_emb(xk, freqs_cis) + xq = comfy.quant_ops.ck.apply_rope_split_half1(xq, freqs_cis) + xk = comfy.quant_ops.ck.apply_rope_split_half1(xk, freqs_cis) present_key_value = None - if past_key_value is not None: + fixed_cache = past_key_value if isinstance(past_key_value, FixedKV) else None + if fixed_cache is not None: + if seq_length == 1 and fixed_cache.index > 0: + # CUDA-graphable decode: write at the device-side ring/linear position + position = fixed_cache.position.view(batch_size, 1, 1, 1).expand_as(xk) + fixed_cache.key.scatter_(2, position, xk) + fixed_cache.value.scatter_(2, position, xv) + output = self._decode_attention(xq, fixed_cache, attention_mask) + return self.o_proj(output), fixed_cache, None + + # prefill: attend the local sequence, persist the tail into the cache + capacity = fixed_cache.key.shape[2] + index = fixed_cache.index + if index + seq_length <= capacity: + fixed_cache.key[:, :, index:index + seq_length] = xk + fixed_cache.value[:, :, index:index + seq_length] = xv + if index > 0: + xk = fixed_cache.key[:, :, :index + seq_length] + xv = fixed_cache.value[:, :, :index + seq_length] + elif index == 0: + # prefill longer than the sliding ring: attend the full local K/V + # (per-query windows come from the prefill sliding mask), cache only + # the last `capacity` keys at their wrapped slots (position % capacity) + slots = torch.arange(seq_length - capacity, seq_length, device=xk.device) % capacity + fixed_cache.key.index_copy_(2, slots, xk[:, :, -capacity:]) + fixed_cache.value.index_copy_(2, slots, xv[:, :, -capacity:]) + else: + raise RuntimeError("gemma4: chunked prefill past the sliding window is not supported") + present_key_value = fixed_cache + elif past_key_value is not None: cumulative_len = 0 if len(past_key_value) > 0: past_key, past_value, cumulative_len = past_key_value @@ -245,6 +321,7 @@ class TransformerBlockGemma4(nn.Module): self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype)) def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None): + output = x sliding_window = None if self.sliding_attention: sliding_window = self.sliding_attention @@ -281,7 +358,8 @@ class TransformerBlockGemma4(nn.Module): x = self.post_per_layer_input_norm(x) x = residual + x - x = x * comfy.ops.cast_to_input(self.layer_scalar, x) + # in-place into the input buffer so CUDA-graph replays land in the static x + x = torch.mul(x, comfy.ops.cast_to_input(self.layer_scalar, x), out=output) return x, present_key_value, shareable_kv @@ -290,6 +368,9 @@ class Gemma4Transformer(nn.Module): def __init__(self, config, device=None, dtype=None, ops=None): super().__init__() self.config = config + self.fixed_kv = True + self.prefetch_dynamic_vbars = True + self.graph_dynamic_vbar_blocks = True self.embed_tokens = _make_scaled_embedding(ops, config.vocab_size, config.hidden_size, config.hidden_size ** 0.5, device, dtype) @@ -298,6 +379,19 @@ class Gemma4Transformer(nn.Module): for i in range(config.num_hidden_layers) ]) + # KV-shared layers never run k_proj/v_proj/k_norm: their never-resolved vbar + # signatures would block layer graph capture, so prefetch only what executes + first_kv_shared = config.num_hidden_layers - config.num_kv_shared_layers if config.num_kv_shared_layers > 0 else config.num_hidden_layers + self._prefetch_units = [] + for i, layer in enumerate(self.layers): + if i >= first_kv_shared: + dead = {layer.self_attn.k_proj, layer.self_attn.v_proj, layer.self_attn.k_norm} + self._prefetch_units.append([ + m for m in layer.modules() if next(m.children(), None) is None and m not in dead + ]) + else: + self._prefetch_units.append(layer) + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype) if config.final_norm else None # Precompute RoPE inv_freq on CPU to match reference code's exact value @@ -311,6 +405,9 @@ class Gemma4Transformer(nn.Module): sliding_inv = 1.0 / (config.rope_theta[1] ** (torch.arange(0, config.head_dim, 2).float() / config.head_dim)) self.register_buffer("_sliding_inv_freq", sliding_inv, persistent=False) + if config.suppress_tokens: + self.register_buffer("_suppress_tokens", torch.tensor(config.suppress_tokens, dtype=torch.long), persistent=False) + # Per-layer input mechanism self.hidden_size_per_layer_input = config.hidden_size_per_layer_input if self.hidden_size_per_layer_input: @@ -322,19 +419,26 @@ class Gemma4Transformer(nn.Module): self.hidden_size_per_layer_input, eps=config.rms_norm_eps, device=device, dtype=dtype) + def get_dynamic_vram__units(self): + return (list(self.layers), []) if self.graph_dynamic_vbar_blocks else ([], []) + def get_past_len(self, past_key_values): for kv in past_key_values: + if isinstance(kv, FixedKV): + return kv.index if len(kv) >= 3: return kv[2] return 0 def _freqs_from_inv(self, inv_freq, position_ids, device, dtype): - """Compute cos/sin from stored inv_freq""" + """Compute per-pair 2x2 rotation matrices [B, 1, S, d/2, 2, 2] from stored inv_freq""" inv_exp = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(device) pos_exp = position_ids[:, None, :].float() freqs = (inv_exp @ pos_exp).transpose(1, 2) - emb = torch.cat((freqs, freqs), dim=-1) - return emb.cos().unsqueeze(1).to(dtype), emb.sin().unsqueeze(1).to(dtype) + cos, sin = freqs.cos(), freqs.sin() + mat = torch.stack((torch.stack((cos, -sin), dim=-1), + torch.stack((sin, cos), dim=-1)), dim=-2) + return mat.unsqueeze(1).to(dtype) def compute_freqs_cis(self, position_ids, device, dtype=None): global_freqs = self._freqs_from_inv(self._global_inv_freq, position_ids, device, dtype) @@ -401,6 +505,72 @@ class Gemma4Transformer(nn.Module): first_kv_shared = self.config.num_hidden_layers - num_kv_shared if num_kv_shared > 0 else self.config.num_hidden_layers shared_sliding_kv = None # KV from last non-shared sliding layer shared_global_kv = None # KV from last non-shared global layer + share_source = {} + if num_kv_shared > 0: + for i in range(first_kv_shared): + share_source[bool(self.layers[i].sliding_attention)] = i + + prefetch_queue = comfy.model_prefetch.make_prefetch_queue( + list(self._prefetch_units), x.device, + {"prefetch_dynamic_vbars": self.prefetch_dynamic_vbars and past_key_values is not None}) + + fixed_kv = (past_key_values is not None and len(past_key_values) > 0 + and isinstance(past_key_values[0], FixedKV)) + decode = fixed_kv and past_len > 0 and seq_len == 1 + # mirror the conditions under which prefetch_queue_pop can actually capture, so + # eager fallbacks keep the sliced decode path instead of the full-capacity one + enable_graph = (decode and mask is None and self.graph_dynamic_vbar_blocks + and prefetch_queue is not None + and hasattr(self.layers[0], "_v_block") + and not comfy.model_management.args.disable_cuda_graphs + and comfy.model_management.is_device_cuda(x.device)) + decode_bias = None + decode_masks = None + if fixed_kv: + prepared = set() + for kv in past_key_values: + if isinstance(kv, FixedKV) and id(kv.position) not in prepared: + kv.prepare(seq_len) + prepared.add(id(kv.position)) + if decode: + if mask is not None: + decode_masks = {} + for kv in past_key_values: + if isinstance(kv, FixedKV) and id(kv.position) not in decode_masks: + decode_masks[id(kv.position)] = _fixed_kv_decode_mask(mask, kv, min_val) + if enable_graph: + # static buffers + per-capacity attention biases: layer graphs replay against + # stable storage, refreshed eagerly each step + capacities = tuple(sorted({kv.key.shape[2] for kv in past_key_values if isinstance(kv, FixedKV)})) + state_key = (x.shape, x.dtype, x.device, tuple(t.shape for t in freqs_cis), capacities, + None if per_layer_inputs is None else per_layer_inputs.shape) + state = getattr(self, "_comfy_cross_step_state", None) + if state is None or state["key"] != state_key: + state = {"key": state_key, + "x": torch.empty_like(x), + "freqs_cis": [torch.empty_like(t) for t in freqs_cis], + "bias": {c: torch.empty((1, 1, 1, c), dtype=x.dtype, device=x.device) for c in capacities}, + "per_layer": None if per_layer_inputs is None else torch.empty_like(per_layer_inputs), + "bias_valid": -1} + self._comfy_cross_step_state = state + comfy.model_management._register_cross_step(self) + state["x"].copy_(x) + for source, target in zip(freqs_cis, state["freqs_cis"]): + target.copy_(source) + x = state["x"] + freqs_cis = state["freqs_cis"] + if per_layer_inputs is not None: + state["per_layer"].copy_(per_layer_inputs) + per_layer_inputs = state["per_layer"] + valid = past_len + 1 + for capacity, bias in state["bias"].items(): + if state["bias_valid"] != past_len: + bias.fill_(min_val) + bias[..., :min(valid, capacity)] = 0 + elif past_len < capacity: + bias[..., past_len:valid] = 0 + state["bias_valid"] = valid + decode_bias = state["bias"] intermediate = None all_intermediate = None @@ -429,12 +599,36 @@ class Gemma4Transformer(nn.Module): is_sliding = hasattr(layer, 'sliding_attention') and layer.sliding_attention if i >= first_kv_shared and num_kv_shared > 0: - shared = shared_sliding_kv if is_sliding else shared_global_kv - if shared is not None: - layer_kwargs['shared_kv'] = shared + if decode: + layer_kwargs['shared_kv'] = past_key_values[share_source[bool(is_sliding)]] + else: + shared = shared_sliding_kv if is_sliding else shared_global_kv + if shared is not None: + layer_kwargs['shared_kv'] = shared - x, current_kv, shareable_kv = layer(x=x, attention_mask=mask, freqs_cis=freqs_cis, past_key_value=past_kv, **layer_kwargs) + if enable_graph: + bias_cache = layer_kwargs.get('shared_kv', past_kv) + layer_mask = decode_bias[bias_cache.key.shape[2]] + elif decode: + bias_cache = layer_kwargs.get('shared_kv', past_kv) + layer_mask = None if decode_masks is None else decode_masks[id(bias_cache.position)] + else: + layer_mask = mask + result = [] + + def core(): + nonlocal x + x, current_kv, shareable_kv = layer(x=x, attention_mask=layer_mask, freqs_cis=freqs_cis, past_key_value=past_kv, **layer_kwargs) + result.append((current_kv, shareable_kv)) + + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, layer, x.dtype, core=core, enable_graph=enable_graph) + + if result: + current_kv, shareable_kv = result[0] + else: + # graph replay: the cache already holds this step's write + current_kv, shareable_kv = past_kv, None next_key_values.append(current_kv if current_kv is not None else ()) # Only track the last sliding/global before the sharing boundary @@ -447,6 +641,14 @@ class Gemma4Transformer(nn.Module): if i == intermediate_output: intermediate = x.clone() + if prefetch_queue is not None: + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, None) + + if fixed_kv: + for kv in past_key_values: + if isinstance(kv, FixedKV): + kv.advance(seq_len) + if self.norm is not None: x = self.norm(x) @@ -481,14 +683,37 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module): if cap: logits = cap * torch.tanh(logits / cap) if self.model.config.suppress_tokens: - logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min + logits.index_fill_(-1, self.model._suppress_tokens, torch.finfo(logits.dtype).min) return logits def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): - past_key_values = [] - for _ in range(self.model.config.num_hidden_layers): - past_key_values.append(()) - return past_key_values + cfg = self.model.config + num_layers = cfg.num_hidden_layers + if not self.model.fixed_kv: + return [() for _ in range(num_layers)] + first_shared = num_layers - cfg.num_kv_shared_layers if cfg.num_kv_shared_layers > 0 else num_layers + # position/seqlen device tensors are shared per cache geometry and filled once per step + trackers = {} + caches = [] + for i in range(num_layers): + if i >= first_shared: + caches.append(()) + continue + sliding = cfg.sliding_attention[i % len(cfg.sliding_attention)] if cfg.sliding_attention else False + head_dim = cfg.head_dim if sliding else cfg.global_head_dim + k_eq_v = cfg.attention_k_eq_v and not sliding + kv_heads = cfg.num_global_key_value_heads if k_eq_v else cfg.num_key_value_heads + length = min(sliding, max_cache_len) if sliding else max_cache_len + cache_cls = RingKV if sliding else FixedKV + tracker = trackers.get((cache_cls, length)) + if tracker is None: + tracker = (torch.empty((batch,), device=device, dtype=torch.int64), + torch.zeros((batch,), device=device, dtype=torch.int32)) + trackers[(cache_cls, length)] = tracker + # zero-init: decode attends full capacity with masked tails, 0*0 stays finite + key = torch.zeros((batch, kv_heads, length, head_dim), device=device, dtype=execution_dtype) + caches.append(cache_cls(key, torch.zeros_like(key), 0, tracker[0], tracker[1])) + return caches def preprocess_embed(self, embed, device): if embed["type"] == "image": @@ -1183,6 +1408,7 @@ def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, po class Gemma4_Tokenizer(): tokenizer_json_data = None + prime_empty_thought = False def state_dict(self): if self.tokenizer_json_data is not None: @@ -1333,8 +1559,8 @@ class Gemma4_Tokenizer(): num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1] n_audio_tokens = self._audio_token_count(num_samples) media += "<|audio>" + "<|audio|>" * n_audio_tokens + "" - # Non-thinking mode primes an empty thought channel so the model answers directly. - model_open = "" if thinking else "<|channel>thought\n" + # 12B/31B prime a closed thought block for non-thinking mode, E2B/E4B must not: it cues them into reasoning inline. + model_open = "<|channel>thought\n" if self.prime_empty_thought and not thinking else "" llama_text = f"{system}<|turn>user\n{text}{media}\n<|turn>model\n{model_open}" text_tokens = super().tokenize_with_weights(llama_text, return_word_ids) @@ -1401,11 +1627,13 @@ class Gemma4SDTokenizer(Gemma4_Tokenizer, sd1_clip.SDTokenizer): def decode(self, token_ids, **kwargs): text = super().decode(token_ids, skip_special_tokens=False) - # Translate thinking channel markers to standard / tags + # Only a close that ends a thought channel becomes : generation primed with + # another channel leaves its opener in the prompt, so its close is not reasoning. + text = re.sub(r"<\|channel>thought\n(.*?)", r"\n\1", text, flags=re.DOTALL) text = text.replace("<|channel>thought\n", "\n") - text = text.replace("", "") # Strip remaining special tokens - text = text.replace("", "").replace("", "").strip() + text = re.sub(r"<\|channel>\w*\n?||<\|turn>\w*\n?|", "", text) + text = text.replace("", "").strip() return text @@ -1418,6 +1646,7 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer): class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer): """Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram.""" embedding_size = 3840 + prime_empty_thought = True def _extract_audio_features(self, waveform, sample_rate): audio = self._resample_16k(waveform, sample_rate) @@ -1443,14 +1672,14 @@ class Gemma4UnifiedTokenizer(Gemma4Tokenizer): class Gemma4Model(sd1_clip.SDClipModel): model_class = None def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): + llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata self.dtypes = set() self.dtypes.add(dtype) super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=self.model_class, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) - def process_tokens(self, tokens, device): - embeds, _, _, _ = super().process_tokens(tokens, device) - return embeds - def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty=0.0): if isinstance(tokens, dict): tokens = next(iter(tokens.values())) @@ -1474,8 +1703,19 @@ class Gemma4Model(sd1_clip.SDClipModel): return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info) +def gemma4_clip_model(model_class): + return type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + + +def gemma4_text_encoder_model(model_class): + return type('Gemma4TextEncoderModel_', (Gemma4Model,), { + 'model_class': model_class, + 'process_tokens': sd1_clip.SDClipModel.process_tokens, + }) + + def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): - clip_model = type('Gemma4Model_', (Gemma4Model,), {'model_class': model_class}) + clip_model = gemma4_clip_model(model_class) class Gemma4TEModel_(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: @@ -1484,12 +1724,15 @@ def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=No if dtype_llama is not None: dtype = dtype_llama super().__init__(device=device, dtype=dtype, name="gemma4", clip_model=clip_model, model_options=model_options) + + def get_dynamic_vram__units(self): + return getattr(self, self.clip).transformer.model.get_dynamic_vram__units() return Gemma4TEModel_ # Variants -def _make_variant(config_cls): +def _make_variant(config_cls, prime_empty_thought=False): audio = config_cls.audio_config is not None bases = (Gemma4AudioMixin, Gemma4Base) if audio else (Gemma4Base,) class Variant(*bases): @@ -1499,8 +1742,8 @@ def _make_variant(config_cls): if audio: self._init_audio(self.model.config, dtype, device, operations) embedding_size = config_cls.hidden_size - if embedding_size != Gemma4SDTokenizer.embedding_size: - tok_cls = type('T', (Gemma4SDTokenizer,), {'embedding_size': embedding_size}) + if embedding_size != Gemma4SDTokenizer.embedding_size or prime_empty_thought: + tok_cls = type('T', (Gemma4SDTokenizer,), {'embedding_size': embedding_size, 'prime_empty_thought': prime_empty_thought}) class Tokenizer(Gemma4Tokenizer): tokenizer_class = tok_cls Variant.tokenizer = Tokenizer @@ -1510,7 +1753,7 @@ def _make_variant(config_cls): Gemma4_E4B = _make_variant(Gemma4Config) Gemma4_E2B = _make_variant(Gemma4_E2B_Config) -Gemma4_31B = _make_variant(Gemma4_31B_Config) +Gemma4_31B = _make_variant(Gemma4_31B_Config, prime_empty_thought=True) # Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant). diff --git a/comfy/text_encoders/hunyuan_video.py b/comfy/text_encoders/hunyuan_video.py index 2ddb4da60..932a3d49b 100644 --- a/comfy/text_encoders/hunyuan_video.py +++ b/comfy/text_encoders/hunyuan_video.py @@ -2,7 +2,7 @@ from comfy import sd1_clip import comfy.model_management import comfy.text_encoders.llama from .hunyuan_image import HunyuanImageTokenizer -from transformers import LlamaTokenizerFast +from .bpe_tokenizer import LlamaTokenizerFast import torch import os import numbers diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index f5c5597ef..49ba5dfa9 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -5,15 +5,33 @@ from typing import Optional, Any, Tuple import math from tqdm import tqdm import comfy.utils +import comfy_kitchen from comfy.ldm.modules.attention import optimized_attention_for_device import comfy.model_management +import comfy.model_prefetch import comfy.ops import comfy.ldm.common_dit import comfy.clip_model from . import qwen_vl + +@dataclass +class FixedKV: + key: torch.Tensor + value: torch.Tensor + index: int + position: torch.Tensor + seqlen: torch.Tensor + + def prepare(self, num_tokens): + self.position.copy_(self.seqlen) + self.seqlen.add_(num_tokens) + + def advance(self, num_tokens): + self.index += num_tokens + @dataclass class Llama2Config: vocab_size: int = 128320 @@ -249,6 +267,9 @@ class Qwen3_8BConfig: rope_scale = None final_norm: bool = True lm_head: bool = True + fixed_kv: bool = False + merged_qkv: bool = False + merged_mlp: bool = False stop_tokens = [151643, 151645] @dataclass @@ -498,9 +519,14 @@ class Attention(nn.Module): self.inner_size = self.num_heads * self.head_dim ops = ops or nn - self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype) - self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=config.qkv_bias, device=device, dtype=dtype) - self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=config.qkv_bias, device=device, dtype=dtype) + self.kv_size = self.num_kv_heads * self.head_dim + self.merged_qkv = getattr(config, "merged_qkv", False) + if self.merged_qkv: + self.qkv_proj = ops.Linear(config.hidden_size, self.inner_size + self.kv_size * 2, bias=config.qkv_bias, device=device, dtype=dtype) + else: + self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype) + self.k_proj = ops.Linear(config.hidden_size, self.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) + self.v_proj = ops.Linear(config.hidden_size, self.kv_size, bias=config.qkv_bias, device=device, dtype=dtype) self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype) self.q_norm = None @@ -522,9 +548,12 @@ class Attention(nn.Module): ): batch_size, seq_length, _ = hidden_states.shape - xq = self.q_proj(hidden_states) - xk = self.k_proj(hidden_states) - xv = self.v_proj(hidden_states) + if self.merged_qkv: + xq, xk, xv = self.qkv_proj(hidden_states).split((self.inner_size, self.kv_size, self.kv_size), dim=-1) + else: + xq = self.q_proj(hidden_states) + xk = self.k_proj(hidden_states) + xv = self.v_proj(hidden_states) xq = xq.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2) xk = xk.view(batch_size, seq_length, self.num_kv_heads, self.head_dim).transpose(1, 2) @@ -537,8 +566,37 @@ class Attention(nn.Module): xq, xk = apply_rope(xq, xk, freqs_cis=freqs_cis) - present_key_value = None - if past_key_value is not None: + fixed_cache = past_key_value if isinstance(past_key_value, FixedKV) else None + if fixed_cache is not None: + xq = xq.transpose(1, 2) + xk = xk.transpose(1, 2) + xv = xv.transpose(1, 2) + if seq_length == 1 and fixed_cache.index > 0: + # CUDA-graphable decode path. + position = fixed_cache.position.view(batch_size, 1, 1, 1).expand_as(xk) + fixed_cache.key.scatter_(1, position, xk) + fixed_cache.value.scatter_(1, position, xv) + output = comfy_kitchen.flash_attention_decode(xq, fixed_cache.key, fixed_cache.value, fixed_cache.seqlen) + return self.o_proj(output.view(batch_size, seq_length, self.inner_size)), fixed_cache + + if attention_mask is None or attention_mask.ndim < 4: + fixed_cache.key[:, :seq_length].copy_(xk) + fixed_cache.value[:, :seq_length].copy_(xv) + else: + valid = attention_mask[:, 0, -1, -seq_length:] == 0 + indices = torch.arange(seq_length, device=xk.device).expand(batch_size, -1) + indices = indices.masked_fill(~valid, seq_length).sort(dim=1).values.clamp_max_(seq_length - 1) + indices = indices.view(batch_size, seq_length, 1, 1).expand_as(xk) + fixed_cache.key[:, :seq_length].copy_(xk.gather(1, indices)) + fixed_cache.value[:, :seq_length].copy_(xv.gather(1, indices)) + fixed_cache.seqlen.copy_(valid.sum(dim=1)) + + xq = xq.transpose(1, 2) + xk = xk.transpose(1, 2) + xv = xv.transpose(1, 2) + + present_key_value = fixed_cache + if fixed_cache is None and past_key_value is not None: index = 0 num_tokens = xk.shape[2] if len(past_key_value) > 0: @@ -569,15 +627,27 @@ class MLP(nn.Module): def __init__(self, config: Llama2Config, device=None, dtype=None, ops: Any = None, intermediate_size=None): super().__init__() intermediate_size = intermediate_size or config.intermediate_size - self.gate_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) - self.up_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) + self.merged_mlp = getattr(config, "merged_mlp", False) + if self.merged_mlp: + self.gate_up_proj = ops.Linear(config.hidden_size, intermediate_size * 2, bias=False, device=device, dtype=dtype) + else: + self.gate_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) + self.up_proj = ops.Linear(config.hidden_size, intermediate_size, bias=False, device=device, dtype=dtype) self.down_proj = ops.Linear(intermediate_size, config.hidden_size, bias=False, device=device, dtype=dtype) if config.mlp_activation == "silu": self.activation = torch.nn.functional.silu + self.merged_input_act = "swiglu" elif config.mlp_activation == "gelu_pytorch_tanh": self.activation = lambda a: torch.nn.functional.gelu(a, approximate="tanh") + self.merged_input_act = None def forward(self, x): + if self.merged_mlp: + x = self.gate_up_proj(x) + if self.merged_input_act is not None: + return comfy.ops.linear_input_act(self.down_proj, x, self.merged_input_act) + gate, up = x.chunk(2, dim=-1) + return self.down_proj(self.activation(gate) * up) return self.down_proj(self.activation(self.gate_proj(x)) * self.up_proj(x)) class TransformerBlock(nn.Module): @@ -596,6 +666,7 @@ class TransformerBlock(nn.Module): optimized_attention=None, past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ): + output = x # Self Attention residual = x x = self.input_layernorm(x) @@ -612,7 +683,7 @@ class TransformerBlock(nn.Module): residual = x x = self.post_attention_layernorm(x) x = self.mlp(x) - x = residual + x + x = torch.add(residual, x, out=output) return x, present_key_value @@ -641,6 +712,7 @@ class TransformerBlockGemma2(nn.Module): optimized_attention=None, past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ): + output = x sliding_window = None if self.transformer_type == 'gemma3': if self.sliding_attention: @@ -676,7 +748,7 @@ class TransformerBlockGemma2(nn.Module): x = self.pre_feedforward_layernorm(x) x = self.mlp(x) x = self.post_feedforward_layernorm(x) - x = residual + x + x = torch.add(residual, x, out=output) return x, present_key_value @@ -688,9 +760,14 @@ def _make_scaled_embedding(ops, vocab_size, hidden_size, scale, device, dtype): class Llama2_(nn.Module): + fixed_kv = False + graph_dynamic_vbar_blocks = False + def __init__(self, config, device=None, dtype=None, ops=None): super().__init__() self.config = config + self.fixed_kv = getattr(config, "fixed_kv", False) + self.graph_dynamic_vbar_blocks = False self.vocab_size = config.vocab_size if self.config.transformer_type == "gemma2" or self.config.transformer_type == "gemma3": @@ -713,8 +790,27 @@ class Llama2_(nn.Module): if config.lm_head: self.lm_head = ops.Linear(config.hidden_size, config.vocab_size, bias=False, device=device, dtype=dtype) + def get_dynamic_vram__units(self): + return (list(self.layers), []) if self.graph_dynamic_vbar_blocks else ([], []) + def get_past_len(self, past_key_values): - return past_key_values[0][2] + first = past_key_values[0] + return first.index if isinstance(first, FixedKV) else first[2] + + def init_kv_cache(self, batch, capacity, device, dtype): + caches = [] + fixed_kv = self.fixed_kv and comfy_kitchen.flash_attention_decode_is_available(device) + for _ in range(self.config.num_hidden_layers): + if fixed_kv: + key = torch.empty((batch, capacity, self.config.num_key_value_heads, self.config.head_dim), device=device, dtype=dtype) + value = torch.empty_like(key) + position = torch.empty((batch,), device=device, dtype=torch.int64) + seqlen = torch.zeros((batch,), device=device, dtype=torch.int32) + caches.append(FixedKV(key, value, 0, position, seqlen)) + else: + key = torch.empty((batch, self.config.num_key_value_heads, capacity, self.config.head_dim), device=device, dtype=dtype) + caches.append((key, torch.empty_like(key), 0)) + return caches def compute_freqs_cis(self, position_ids, device): return precompute_freqs_cis(self.config.head_dim, @@ -736,6 +832,10 @@ class Llama2_(nn.Module): past_len = 0 if past_key_values is not None and len(past_key_values) > 0: past_len = self.get_past_len(past_key_values) + fixed_kv = past_key_values is not None and len(past_key_values) > 0 and isinstance(past_key_values[0], FixedKV) + fixed_kv_decode = fixed_kv and past_len > 0 and seq_len == 1 + if fixed_kv_decode: + attention_mask = None if position_ids is None: position_ids = torch.arange(past_len, past_len + seq_len, device=x.device).unsqueeze(0) @@ -756,6 +856,32 @@ class Llama2_(nn.Module): optimized_attention = optimized_attention_for_device(x.device, mask=mask is not None, small_input=True) + enable_graph = self.graph_dynamic_vbar_blocks and fixed_kv_decode + if enable_graph: + freqs_cis_groups = freqs_cis if isinstance(freqs_cis, list) else [freqs_cis] + cross_step_state_key = [(x.shape, x.stride(), x.dtype, x.device)] + for group in freqs_cis_groups: + for tensor in group: + cross_step_state_key.append((tensor.shape, tensor.stride(), tensor.dtype, tensor.device)) + cross_step_state_key = tuple(cross_step_state_key) + cross_step_state = getattr(self, "_comfy_cross_step_state", None) + if cross_step_state is None or cross_step_state["key"] != cross_step_state_key: + static_freqs_cis = [] + for group in freqs_cis_groups: + static_freqs_cis.append(tuple(torch.empty_like(tensor) for tensor in group)) + if not isinstance(freqs_cis, list): + static_freqs_cis = static_freqs_cis[0] + cross_step_state = {"key": cross_step_state_key, "x": torch.empty_like(x), "freqs_cis": static_freqs_cis} + self._comfy_cross_step_state = cross_step_state + comfy.model_management._register_cross_step(self) + cross_step_state["x"].copy_(x) + static_freqs_cis_groups = cross_step_state["freqs_cis"] if isinstance(freqs_cis, list) else [cross_step_state["freqs_cis"]] + for source_group, target_group in zip(freqs_cis_groups, static_freqs_cis_groups): + for source, target in zip(source_group, target_group): + target.copy_(source) + x = cross_step_state["x"] + freqs_cis = cross_step_state["freqs_cis"] + intermediate = None all_intermediate = None only_layers = None @@ -769,7 +895,8 @@ class Llama2_(nn.Module): elif intermediate_output < 0: intermediate_output = len(self.layers) + intermediate_output - next_key_values = [] + prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.layers), x.device, {"prefetch_dynamic_vbars": getattr(self, "prefetch_dynamic_vbars", False)}) + next_key_values = list(past_key_values) if past_key_values is not None else [] for i, layer in enumerate(self.layers): if all_intermediate is not None: if only_layers is None or (i in only_layers): @@ -779,16 +906,24 @@ class Llama2_(nn.Module): if past_key_values is not None: past_kv = past_key_values[i] if len(past_key_values) > 0 else [] - x, current_kv = layer( - x=x, - attention_mask=mask, - freqs_cis=freqs_cis, - optimized_attention=optimized_attention, - past_key_value=past_kv, - ) + if fixed_kv: + past_kv.prepare(seq_len) - if current_kv is not None: - next_key_values.append(current_kv) + def core(): + nonlocal x + x, current_kv = layer( + x=x, + attention_mask=mask, + freqs_cis=freqs_cis, + optimized_attention=optimized_attention, + past_key_value=past_kv, + ) + if next_key_values: + next_key_values[i] = current_kv + + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, layer, x.dtype, core=core, enable_graph=enable_graph) + if fixed_kv: + next_key_values[i].advance(seq_len) # DeepStack: add per-layer visual features into the first len() decoder layers at image positions (Qwen3-VL) if deepstack_embeds is not None and i < len(deepstack_embeds): @@ -797,6 +932,9 @@ class Llama2_(nn.Module): if i == intermediate_output: intermediate = x.clone() + if prefetch_queue is not None: + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, x.device, None) + if self.norm is not None: x = self.norm(x) @@ -810,7 +948,7 @@ class Llama2_(nn.Module): if intermediate is not None and final_layer_norm_intermediate and self.norm is not None: intermediate = self.norm(intermediate) - if len(next_key_values) > 0: + if next_key_values: return x, intermediate, next_key_values else: return x, intermediate @@ -868,24 +1006,13 @@ class BaseGenerate: else: module = self.model.embed_tokens - offload_stream = None - if module.comfy_cast_weights: - weight, _, offload_stream = comfy.ops.cast_bias_weight(module, input, offloadable=True) - else: - weight = self.model.embed_tokens.weight.to(x) - - x = torch.nn.functional.linear(input, weight, None) - - comfy.ops.uncast_bias_weight(module, weight, None, offload_stream) - return x + if not module.comfy_cast_weights: + return torch.nn.functional.linear(input, self.model.embed_tokens.weight.to(x), None) + with comfy.ops.CastBiasWeightContext(module, input, offloadable=True) as (weight, _bias): + return torch.nn.functional.linear(input, weight, None) def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): - model_config = self.model.config - past_key_values = [] - for x in range(model_config.num_hidden_layers): - past_key_values.append((torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), - torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0)) - return past_key_values + return self.model.init_kv_cache(batch, max_cache_len, device, execution_dtype) def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None): device = embeds.device diff --git a/comfy/text_encoders/lt.py b/comfy/text_encoders/lt.py index bc5cbae28..c512a7d48 100644 --- a/comfy/text_encoders/lt.py +++ b/comfy/text_encoders/lt.py @@ -81,6 +81,17 @@ class LTXAVGemmaTokenizer(sd1_clip.SD1Tokenizer): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma3_12b", tokenizer=Gemma3_12BTokenizer) +def ltxav_gemma4_tokenizer(tokenizer): + class LTXAVGemma4Tokenizer(tokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data) + gemma_tokenizer = getattr(self, self.clip) + if gemma_tokenizer.min_length == 1: + gemma_tokenizer.min_length = 1024 + + return LTXAVGemma4Tokenizer + + class Gemma3_12BModel(sd1_clip.SDClipModel): def __init__(self, device="cpu", layer="all", layer_idx=None, dtype=None, attention_mask=True, model_options={}): llama_quantization_metadata = model_options.get("llama_quantization_metadata", None) @@ -97,10 +108,10 @@ class Gemma3_12BModel(sd1_clip.SDClipModel): return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, stop_tokens=[106], presence_penalty=presence_penalty) # 106 is class DualLinearProjection(torch.nn.Module): - def __init__(self, in_dim, out_dim_video, out_dim_audio, dtype=None, device=None, operations=None): + def __init__(self, in_dim, out_dim_video, out_dim_audio, video_bias=True, audio_bias=True, dtype=None, device=None, operations=None): super().__init__() - self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=True, dtype=dtype, device=device) - self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=True, dtype=dtype, device=device) + self.audio_aggregate_embed = operations.Linear(in_dim, out_dim_audio, bias=audio_bias, dtype=dtype, device=device) + self.video_aggregate_embed = operations.Linear(in_dim, out_dim_video, bias=video_bias, dtype=dtype, device=device) def forward(self, x): source_dim = x.shape[-1] @@ -112,22 +123,28 @@ class DualLinearProjection(torch.nn.Module): return torch.cat((video, audio), dim=-1) class LTXAVTEModel(torch.nn.Module): - def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", model_options={}): + def __init__(self, dtype_llama=None, device="cpu", dtype=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True, model_options={}): super().__init__() self.dtypes = set() self.dtypes.add(dtype) self.compat_mode = False self.text_projection_type = text_projection_type + self.text_encoder_key = text_encoder_key + self.execution_device = None - self.gemma3_12b = Gemma3_12BModel(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) + self.gemma3_12b = text_encoder_model(device=device, dtype=dtype_llama, model_options=model_options, layer="all", layer_idx=None) self.dtypes.add(dtype_llama) operations = self.gemma3_12b.operations # TODO + text_encoder_config = self.gemma3_12b.transformer.model.config + projection_in_dim = text_encoder_config.hidden_size * (text_encoder_config.num_hidden_layers + 1) + if video_projection_bias is None: + video_projection_bias = self.text_projection_type == "dual_linear" if self.text_projection_type == "single_linear": - self.text_embedding_projection = operations.Linear(3840 * 49, 3840, bias=False, dtype=dtype, device=device) + self.text_embedding_projection = operations.Linear(projection_in_dim, video_projection_dim, bias=video_projection_bias, dtype=dtype, device=device) elif self.text_projection_type == "dual_linear": - self.text_embedding_projection = DualLinearProjection(3840 * 49, 4096, 2048, dtype=dtype, device=device, operations=operations) + self.text_embedding_projection = DualLinearProjection(projection_in_dim, video_projection_dim, audio_projection_dim, video_bias=video_projection_bias, audio_bias=audio_projection_bias, dtype=dtype, device=device, operations=operations) def enable_compat_mode(self): # TODO: remove @@ -161,7 +178,7 @@ class LTXAVTEModel(torch.nn.Module): self.execution_device = None def encode_token_weights(self, token_weight_pairs): - token_weight_pairs = token_weight_pairs["gemma3_12b"] + token_weight_pairs = token_weight_pairs[self.text_encoder_key] out, pooled, extra = self.gemma3_12b.encode_token_weights(token_weight_pairs) out = out[:, :, -torch.sum(extra["attention_mask"]).item():] @@ -189,51 +206,54 @@ class LTXAVTEModel(torch.nn.Module): return out.to(device=out_device, dtype=torch.float), pooled, extra def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty): - return self.gemma3_12b.generate(tokens["gemma3_12b"], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) + return self.gemma3_12b.generate(tokens[self.text_encoder_key], do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty) def load_sd(self, sd): - if "model.layers.47.self_attn.q_norm.weight" in sd: - return self.gemma3_12b.load_sd(sd) - else: - sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.weight": "text_embedding_projection.weight", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) - if len(sdo) == 0: - sdo = sd + missing_all = [] + unexpected_all = [] - missing_all = [] - unexpected_all = [] + if "model.layers.0.self_attn.q_norm.weight" in sd: + gemma_sd = {k: v for k, v in sd.items() if not k.startswith("text_embedding_projection.")} + missing, unexpected = self.gemma3_12b.load_sd(gemma_sd) + missing_all.extend(missing) + unexpected_all.extend(unexpected) - for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: - component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} - if component_sd: - missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) - missing_all.extend([f"{prefix}{k}" for k in missing]) - unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) + sdo = comfy.utils.state_dict_prefix_replace(sd, {"text_embedding_projection.aggregate_embed.": "text_embedding_projection.", "text_embedding_projection.": "text_embedding_projection."}, filter_keys=True) + if len(sdo) == 0: + sdo = sd - if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove - ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) - if ww is not None: - if ww.shape[0] == 3840: - self.enable_compat_mode() - sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) - self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) - sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) - self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + for prefix, component in [("text_embedding_projection.", self.text_embedding_projection)]: + component_sd = {k.replace(prefix, ""): v for k, v in sdo.items() if k.startswith(prefix)} + if component_sd: + missing, unexpected = component.load_state_dict(component_sd, strict=False, assign=getattr(self, "can_assign_sd", False)) + missing_all.extend([f"{prefix}{k}" for k in missing]) + unexpected_all.extend([f"{prefix}{k}" for k in unexpected]) - return (missing_all, unexpected_all) + if "model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.2.attn1.to_q.bias" not in sd: # TODO: remove + ww = sd.get("model.diffusion_model.audio_embeddings_connector.transformer_1d_blocks.0.attn1.to_q.bias", None) + if ww is not None: + if ww.shape[0] == 3840: + self.enable_compat_mode() + sdv = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.video_embeddings_connector.": ""}, filter_keys=True) + self.video_embeddings_connector.load_state_dict(sdv, strict=False, assign=getattr(self, "can_assign_sd", False)) + sda = comfy.utils.state_dict_prefix_replace(sd, {"model.diffusion_model.audio_embeddings_connector.": ""}, filter_keys=True) + self.audio_embeddings_connector.load_state_dict(sda, strict=False, assign=getattr(self, "can_assign_sd", False)) + + return (missing_all, unexpected_all) def memory_estimation_function(self, token_weight_pairs, device=None): constant = 6.0 if comfy.model_management.should_use_bf16(device): constant /= 2.0 - token_weight_pairs = token_weight_pairs.get("gemma3_12b", []) + token_weight_pairs = token_weight_pairs.get(self.text_encoder_key, []) m = min([sum(1 for _ in itertools.takewhile(lambda x: x[0] == 0, sub)) for sub in token_weight_pairs]) num_tokens = sum(map(lambda a: len(a), token_weight_pairs)) - m num_tokens = max(num_tokens, 642) return num_tokens * constant * 1024 * 1024 -def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear"): +def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection_type="single_linear", text_encoder_model=Gemma3_12BModel, text_encoder_key="gemma3_12b", video_projection_dim=3840, audio_projection_dim=2048, video_projection_bias=None, audio_projection_bias=True): class LTXAVTEModel_(LTXAVTEModel): def __init__(self, device="cpu", dtype=None, model_options={}): if llama_quantization_metadata is not None: @@ -241,16 +261,29 @@ def ltxav_te(dtype_llama=None, llama_quantization_metadata=None, text_projection model_options["llama_quantization_metadata"] = llama_quantization_metadata if dtype_llama is not None: dtype = dtype_llama - super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, model_options=model_options) + super().__init__(dtype_llama=dtype_llama, device=device, dtype=dtype, text_projection_type=text_projection_type, text_encoder_model=text_encoder_model, text_encoder_key=text_encoder_key, video_projection_dim=video_projection_dim, audio_projection_dim=audio_projection_dim, video_projection_bias=video_projection_bias, audio_projection_bias=audio_projection_bias, model_options=model_options) return LTXAVTEModel_ def sd_detect(state_dict_list, prefix=""): for sd in state_dict_list: - if "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd: - return {"text_projection_type": "dual_linear"} - if "{}text_embedding_projection.weight".format(prefix) in sd or "{}text_embedding_projection.aggregate_embed.weight".format(prefix) in sd: - return {"text_projection_type": "single_linear"} + video_key = "{}text_embedding_projection.video_aggregate_embed.weight".format(prefix) + audio_key = "{}text_embedding_projection.audio_aggregate_embed.weight".format(prefix) + if video_key in sd and audio_key in sd: + return { + "text_projection_type": "dual_linear", + "video_projection_dim": sd[video_key].shape[0], + "audio_projection_dim": sd[audio_key].shape[0], + "video_projection_bias": "{}text_embedding_projection.video_aggregate_embed.bias".format(prefix) in sd, + "audio_projection_bias": "{}text_embedding_projection.audio_aggregate_embed.bias".format(prefix) in sd, + } + for key in ("{}text_embedding_projection.weight".format(prefix), "{}text_embedding_projection.aggregate_embed.weight".format(prefix)): + if key in sd: + return { + "text_projection_type": "single_linear", + "video_projection_dim": sd[key].shape[0], + "video_projection_bias": key.removesuffix("weight") + "bias" in sd, + } return {} diff --git a/comfy/text_encoders/lumina2.py b/comfy/text_encoders/lumina2.py index b1f1dbb9f..e44920203 100644 --- a/comfy/text_encoders/lumina2.py +++ b/comfy/text_encoders/lumina2.py @@ -49,10 +49,6 @@ class Gemma3_4B_Vision_Model(sd1_clip.SDClipModel): super().__init__(device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, dtype=dtype, special_tokens={"start": 2, "pad": 0}, layer_norm_hidden_state=False, model_class=comfy.text_encoders.llama.Gemma3_4B_Vision, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) - def process_tokens(self, tokens, device): - embeds, _, _, _ = super().process_tokens(tokens, device) - return embeds - class LuminaModel(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}, name="gemma2_2b", clip_model=Gemma2_2BModel): super().__init__(device=device, dtype=dtype, name=name, clip_model=clip_model, model_options=model_options) diff --git a/comfy/text_encoders/minimax.py b/comfy/text_encoders/minimax.py index c2dc47f7f..d79ccf0ea 100644 --- a/comfy/text_encoders/minimax.py +++ b/comfy/text_encoders/minimax.py @@ -127,10 +127,6 @@ class MiniMaxH3Tokenizer(comfy.sd1_clip.SD1Tokenizer): tokenizer = lambda *a, **kw: Qwen3VLSDTokenizer(*a, **kw, embedding_size=5120, embedding_key="qwen3vl_32b") super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="qwen3vl_32b", tokenizer=tokenizer) - def _text_ids(self, text): - tok = self.qwen3vl_32b.tokenizer - return tok(text, add_special_tokens=False)["input_ids"] - @staticmethod def _vision_entry(data, video_block=False): emb = {"type": "image", "data": data, "original_type": "image"} @@ -143,7 +139,16 @@ class MiniMaxH3Tokenizer(comfy.sd1_clip.SD1Tokenizer): entries = [] def add_text(s): - entries.extend((tid, 1.0) for tid in self._text_ids(s)) + if not s: + return + token_batches = self.qwen3vl_32b.tokenize_with_weights( + s, + return_word_ids=False, + disable_weights=True, + ) + if len(token_batches) != 1: + raise ValueError("MiniMax H3 text segment exceeds the supported prompt length.") + entries.extend(token_batches[0]) def add_vision(data, video_block=False): entries.append((VISION_START, 1.0)) diff --git a/comfy/text_encoders/minimax_music.py b/comfy/text_encoders/minimax_music.py new file mode 100644 index 000000000..c88d463cc --- /dev/null +++ b/comfy/text_encoders/minimax_music.py @@ -0,0 +1,117 @@ +import torch +from tokenizers import Tokenizer + +import comfy.ops +from comfy.ldm.minimax_music.ar import CFG_SCALE, CFG_TOP_K, MAX_AUDIO_FRAMES, MiniMaxMusic3AR +from comfy.ldm.minimax_music.prompt import SPECIAL_TOKEN_IDS, build_prompt + + +MODEL_CONFIG = { + "vocab_size": 200000, + "hidden_size": 4096, + "intermediate_size": 12288, + "num_hidden_layers": 36, + "num_attention_heads": 32, + "num_key_value_heads": 8, + "max_position_embeddings": 10240, + "rms_norm_eps": 1e-6, + "rope_theta": 1000000.0, + "head_dim": 128, + "audio_vocab_size": 1024, + "audio_num_codebooks": 8, + "decoder_num_heads": 16, + "decoder_intermediate_size": 6144, + "decoder_num_layers": 4, +} + + +def detect_merged_config(state_dict, prefix=""): + return { + "merged_qkv": "{}model.layers.0.self_attn.qkv_proj.weight".format(prefix) in state_dict, + "merged_mlp": "{}model.layers.0.mlp.gate_up_proj.weight".format(prefix) in state_dict, + "decoder_merged_qkv": "{}model.audio_decoder.layers.0.self_attn.qkv_proj.weight".format(prefix) in state_dict, + "decoder_merged_mlp": "{}model.audio_decoder.layers.0.mlp.gate_up_proj.weight".format(prefix) in state_dict, + } + + +class MiniMaxMusic3Tokenizer: + def __init__(self, embedding_directory=None, tokenizer_data={}): + tokenizer_json = tokenizer_data.get("tokenizer_json") + if tokenizer_json is None: + raise ValueError("MiniMax Music3 text encoder checkpoint is missing tokenizer_json") + if torch.is_tensor(tokenizer_json): + tokenizer_json = tokenizer_json.detach().cpu().numpy().tobytes() + self.tokenizer_json = tokenizer_json + self.tokenizer = Tokenizer.from_str(tokenizer_json.decode("utf-8")) + for token, expected in SPECIAL_TOKEN_IDS.items(): + if self.tokenizer.token_to_id(token) != expected: + raise ValueError(f"MiniMax Music3 tokenizer mismatch for {token}") + + def tokenize_with_weights(self, text, return_word_ids=False, **kwargs): + prompt = build_prompt(text, kwargs.get("lyrics", "")) + token_ids = self.tokenizer.encode(prompt, add_special_tokens=False).ids + return { + "minimax_music3": [[(token, 1.0) for token in token_ids]], + "seed": int(kwargs.get("seed", 0)), + "max_audio_frames": int(kwargs.get("max_audio_frames", MAX_AUDIO_FRAMES)), + "cfg_scale": float(kwargs.get("cfg_scale", CFG_SCALE)), + "top_k": int(kwargs.get("top_k", CFG_TOP_K)), + } + + def state_dict(self): + return {"tokenizer_json": torch.frombuffer(bytearray(self.tokenizer_json), dtype=torch.uint8)} + + def decode(self, token_ids, skip_special_tokens=True): + return self.tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens) + + +class MiniMaxMusic3TEModel(MiniMaxMusic3AR): + def __init__(self, device="cpu", dtype=None, model_options={}, projection_config=None): + dtype = torch.bfloat16 + quant_config = model_options.get("quantization_metadata", None) + operations = model_options.get("custom_operations", None) + if operations is None: + operations = comfy.ops.mixed_precision_ops(quant_config, dtype) if quant_config is not None else comfy.ops.manual_cast + super().__init__({**MODEL_CONFIG, **(projection_config or {})}, dtype, device, operations) + self.dtypes = {dtype} + self.execution_device = device + + def set_clip_options(self, options): + self.execution_device = options.get("execution_device", self.execution_device) + + def reset_clip_options(self): + pass + + def get_dynamic_vram__units(self): + units, last_units = self.model.get_dynamic_vram__units() + if self.model.pruned_embedding: + last_units = [*last_units, self.model.embed_tokens_prefill] + return [(self.model.audio_decoder, self.model.audio_extra_embedding), *units], last_units + + def encode_token_weights(self, token_weight_pairs): + token_ids = [token for token, _ in token_weight_pairs["minimax_music3"][0]] + input_ids = torch.tensor([token_ids], dtype=torch.long) + seed = token_weight_pairs["seed"] + max_audio_frames = token_weight_pairs["max_audio_frames"] + cfg_scale = token_weight_pairs["cfg_scale"] + top_k = token_weight_pairs["top_k"] + hidden = self.generate(input_ids, seed, max_audio_frames, self.execution_device, cfg_scale, top_k) + return hidden.unsqueeze(0), None, {} + + def load_state_dict(self, state_dict, strict=True, assign=False): + if self.model.pruned_embedding is None: + self.model.pruned_embedding = "model.embed_tokens_prefill.weight" in state_dict + if self.model.pruned_embedding: + del self.model.embed_tokens + else: + del self.model.embed_tokens_prefill, self.model.embed_tokens_audio + if self.model.pruned_lm_head is None: + self.model.pruned_lm_head = "model.lm_head_pruned.weight" in state_dict + if self.model.pruned_lm_head: + del self.model.lm_head + else: + del self.model.lm_head_pruned + return super().load_state_dict(state_dict, strict=strict, assign=assign) + + def load_sd(self, state_dict): + return self.load_state_dict(state_dict, strict=False, assign=getattr(self, "can_assign_sd", False)) diff --git a/comfy/utils.py b/comfy/utils.py index 61c2a22dd..b31b70daa 100644 --- a/comfy/utils.py +++ b/comfy/utils.py @@ -82,19 +82,47 @@ _TYPES = { "U16": torch.uint16, } +_SAFETENSORS_MAX_HEADER_SIZE = 100_000_000 + + +def _invalid_safetensors_error(message, ckpt): + return ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt or invalid. Make sure this is actually a safetensors file and not a ckpt or pt or other filetype.".format(message, ckpt)) + + +def _incomplete_safetensors_error(message, ckpt): + return ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt/incomplete. Check the file size and make sure you have copied/downloaded it correctly.".format(message, ckpt)) + + def load_safetensors(ckpt): import comfy_aimdo.model_mmap + file_size = os.path.getsize(ckpt) + if file_size < 8: + raise _incomplete_safetensors_error("The safetensors header is incomplete.", ckpt) + file_lock = threading.Lock() model_mmap = comfy_aimdo.model_mmap.ModelMMAP(ckpt) f = model_mmap.get_file_handle() - file_size = os.path.getsize(ckpt) mv = memoryview((ctypes.c_uint8 * file_size).from_address(model_mmap.get())) header_size = struct.unpack(" _SAFETENSORS_MAX_HEADER_SIZE: + raise _invalid_safetensors_error("The safetensors header is too large.", ckpt) - mv = mv[(data_base_offset := 8 + header_size):] + data_base_offset = 8 + header_size + if data_base_offset > file_size: + raise _incomplete_safetensors_error("The safetensors header is incomplete.", ckpt) + + try: + header = json.loads(mv[8:data_base_offset].tobytes().decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError) as e: + raise _invalid_safetensors_error(str(e), ckpt) from e + + if not isinstance(header, dict): + raise _invalid_safetensors_error("The safetensors header is invalid.", ckpt) + + mv = mv[data_base_offset:] + data_size = len(mv) sd = {} for name, info in header.items(): @@ -102,13 +130,21 @@ def load_safetensors(ckpt): continue start, end = info["data_offsets"] + dtype = _TYPES[info["dtype"]] + if start < 0 or end < start: + raise _invalid_safetensors_error("Tensor '{}' has invalid data offsets.".format(name), ckpt) + if end > data_size: + raise _incomplete_safetensors_error("Tensor '{}' extends past the end of the file.".format(name), ckpt) + if math.prod(info["shape"]) * dtype.itemsize != end - start: + raise _invalid_safetensors_error("Tensor '{}' does not match its declared shape and dtype.".format(name), ckpt) + if start == end: - sd[name] = torch.empty(info["shape"], dtype =_TYPES[info["dtype"]]) + sd[name] = torch.empty(info["shape"], dtype=dtype) else: with warnings.catch_warnings(): #We are working with read-only RAM by design warnings.filterwarnings("ignore", message="The given buffer is not writable") - tensor = torch.frombuffer(mv[start:end], dtype=_TYPES[info["dtype"]]).view(info["shape"]) + tensor = torch.frombuffer(mv[start:end], dtype=dtype).view(info["shape"]) storage = tensor.untyped_storage() setattr(storage, "_comfy_tensor_file_slice", @@ -143,9 +179,9 @@ def load_torch_file(ckpt, safe_load=False, device=None, return_metadata=False): if len(e.args) > 0: message = e.args[0] if "HeaderTooLarge" in message: - raise ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt or invalid. Make sure this is actually a safetensors file and not a ckpt or pt or other filetype.".format(message, ckpt)) + raise _invalid_safetensors_error(message, ckpt) if "MetadataIncompleteBuffer" in message: - raise ValueError("{}\n\nFile path: {}\n\nThe safetensors file is corrupt/incomplete. Check the file size and make sure you have copied/downloaded it correctly.".format(message, ckpt)) + raise _incomplete_safetensors_error(message, ckpt) raise e else: torch_args = {} diff --git a/comfy_api/latest/_input/video_types.py b/comfy_api/latest/_input/video_types.py index c9c153e06..1f2a0b867 100644 --- a/comfy_api/latest/_input/video_types.py +++ b/comfy_api/latest/_input/video_types.py @@ -30,15 +30,24 @@ class VideoInput(ABC): metadata: Optional[dict] = None, bit_depth: int | None = None, crf: float | None = None, + color_space: str | None = None, ): """ Abstract method to save the video input to a file. bit_depth selects the encoded bit depth; None keeps the video's native depth. - crf selects the H.264 constant rate factor; None uses the encoder default. + crf selects the H.264 or AV1 constant rate factor; None uses the encoder default. + color_space="sRGB" writes SDR BT.709/sRGB video. "HDR" writes 10-bit BT.2020/HLG video; + "HDR PQ" selects BT.2020/PQ. + Tensor-created videos default to sRGB when color_space is None. Loaded videos keep matching recognized native color + properties; other input pixels must already use the selected color space. """ pass + def get_color_space(self) -> str: + """Return the video's color space as sRGB, HDR, HDR PQ, or auto when unspecified.""" + return "auto" + @abstractmethod def as_trimmed( self, diff --git a/comfy_api/latest/_input_impl/video_types.py b/comfy_api/latest/_input_impl/video_types.py index 5f99e61fb..b5fe41895 100644 --- a/comfy_api/latest/_input_impl/video_types.py +++ b/comfy_api/latest/_input_impl/video_types.py @@ -1,6 +1,6 @@ from av.container import InputContainer from av.subtitles.stream import SubtitleStream -from av.video.reformatter import ColorRange +from av.video.reformatter import ColorPrimaries, ColorRange, ColorTrc from fractions import Fraction from typing import Optional from .._input import AudioInput, VideoInput @@ -16,6 +16,38 @@ from .._util import VideoContainer, VideoCodec, VideoComponents, normalize_crop_ import logging +VIDEO_ENCODERS = { + VideoCodec.H264: "h264", + VideoCodec.AV1: "libsvtav1", +} +VIDEO_CONTAINER_FORMATS = { + VideoContainer.MP4: "mp4", + VideoContainer.MKV: "matroska", + VideoContainer.WEBM: "webm", +} +WEBM_STREAM_CODECS = { + "video": {"av1", "vp8", "vp9"}, + "audio": {"opus", "vorbis"}, + "subtitle": {"webvtt"}, +} +BT2020_NCL = 9 +BT709_NCL = 1 +HDR_COLOR_TRANSFERS = { + "HDR": ColorTrc.ARIB_STD_B67, + "HDR PQ": ColorTrc.SMPTE2084, +} +VIDEO_COLOR_TRANSFERS = { + "sRGB": ColorTrc.IEC61966_2_1, + **HDR_COLOR_TRANSFERS, +} +VIDEO_TRANSFER_COLOR_SPACES = { + ColorTrc.BT709: "sRGB", + ColorTrc.IEC61966_2_1: "sRGB", + ColorTrc.ARIB_STD_B67: "HDR", + ColorTrc.SMPTE2084: "HDR PQ", +} + + def container_to_output_format(container_format: str | None) -> str | None: """ A container's `format` may be a comma-separated list of formats. @@ -37,22 +69,24 @@ def get_open_write_kwargs( ) -> dict: """Get kwargs for writing a `VideoFromFile` to a file/stream with `av.open`""" is_write_to_buffer = isinstance(dest, io.BytesIO) - is_mp4_file = not is_write_to_buffer and os.path.splitext(dest)[1].lower() == ".mp4" - movflags = "use_metadata_tags+faststart" if is_mp4_file else "use_metadata_tags" - open_kwargs = { - "mode": "w", - # If isobmff, preserve custom metadata tags (workflow, prompt, extra_pnginfo) - "options": {"movflags": movflags}, - } + open_kwargs = {"mode": "w"} if is_write_to_buffer: # Set output format explicitly, since it cannot be inferred from file extension if to_format == VideoContainer.AUTO: to_format = container_format.lower() + elif isinstance(to_format, VideoContainer): + to_format = VIDEO_CONTAINER_FORMATS[to_format] elif isinstance(to_format, str): to_format = to_format.lower() open_kwargs["format"] = container_to_output_format(to_format) + output_format = open_kwargs["format"] if is_write_to_buffer else os.path.splitext(dest)[1].lower().lstrip(".") + if output_format in ("mov", "mp4"): + # Preserve custom metadata tags (workflow, prompt, extra_pnginfo) in isobmff. + movflags = "use_metadata_tags" if is_write_to_buffer else "use_metadata_tags+faststart" + open_kwargs["options"] = {"movflags": movflags} + return open_kwargs @@ -100,19 +134,66 @@ def write_output_metadata(container: InputContainer, output, metadata: dict | No output.metadata[key] = value if isinstance(value, str) else json.dumps(value) -def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> dict: - if format != VideoContainer.AUTO and format != VideoContainer.MP4: - raise ValueError("Only MP4 format is supported for now") - if codec != VideoCodec.AUTO and codec != VideoCodec.H264: - raise ValueError("Only H264 codec is supported for now") +def video_output_config(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> tuple[dict, VideoContainer, VideoCodec]: + if isinstance(format, str): + format = VideoContainer(format) + if isinstance(codec, str): + codec = VideoCodec(codec) + + if format == VideoContainer.AUTO: + extension = os.path.splitext(os.fspath(path))[1].lower() if isinstance(path, (str, os.PathLike)) else "" + format = { + ".mkv": VideoContainer.MKV, + ".webm": VideoContainer.WEBM, + }.get(extension, VideoContainer.MP4) + if codec == VideoCodec.AUTO: + codec = VideoCodec.AV1 if format == VideoContainer.WEBM else VideoCodec.H264 + if format == VideoContainer.WEBM and codec != VideoCodec.AV1: + raise ValueError("WebM output requires the AV1 codec") + # FFmpeg's faststart pass reopens the output by filename, so it cannot be used with file-like objects. - movflags = "use_metadata_tags+faststart" if isinstance(path, (str, os.PathLike)) else "use_metadata_tags" - open_kwargs = {"mode": "w", "options": {"movflags": movflags}} - if isinstance(format, VideoContainer) and format != VideoContainer.AUTO: - open_kwargs["format"] = format.value - elif isinstance(path, io.BytesIO): - open_kwargs["format"] = "mp4" # no file extension to infer the format from - return open_kwargs + open_kwargs = {"mode": "w", "format": VIDEO_CONTAINER_FORMATS[format]} + if format == VideoContainer.MP4: + movflags = "use_metadata_tags+faststart" if isinstance(path, (str, os.PathLike)) else "use_metadata_tags" + open_kwargs["options"] = {"movflags": movflags} + return open_kwargs, format, codec + + +def set_video_color_properties(target, color_space): + is_hdr = color_space in HDR_COLOR_TRANSFERS + target.color_primaries = ColorPrimaries.BT2020 if is_hdr else ColorPrimaries.BT709 + target.color_trc = VIDEO_COLOR_TRANSFERS[color_space] + target.colorspace = BT2020_NCL if is_hdr else BT709_NCL + target.color_range = ColorRange.MPEG + + +def copy_color_properties(source, target): + target.color_primaries = source.color_primaries + target.color_trc = source.color_trc + target.colorspace = source.colorspace + target.color_range = source.color_range + + +def video_stream_color_space(stream) -> str | None: + if stream is None: + return None + return VIDEO_TRANSFER_COLOR_SPACES.get(stream.color_trc) + + +def video_encoder_options(codec: VideoCodec, crf: float | None) -> dict[str, str]: + if crf is None: + return {} + if codec == VideoCodec.AV1 and crf == 0: + return {"svtav1-params": "lossless=1"} + return {"crf": str(crf)} + + +def webm_streams_compatible(streams) -> bool: + for stream in streams: + allowed_codecs = WEBM_STREAM_CODECS.get(stream.type) + if allowed_codecs is not None and stream.codec_context is not None and stream.codec.canonical_name not in allowed_codecs: + return False + return True def _rotation_quadrant(frame: av.VideoFrame) -> int: @@ -197,6 +278,13 @@ class VideoFromFile(VideoInput): video_stream = container.streams.video[0] if len(container.streams.video) > 0 else None return video_stream_bit_depth(video_stream) + def get_color_space(self) -> str: + if isinstance(self.__file, io.BytesIO): + self.__file.seek(0) + with av.open(self.__file, mode="r") as container: + video_stream = container.streams.video[0] if len(container.streams.video) > 0 else None + return video_stream_color_space(video_stream) or "sRGB" + def get_duration(self) -> float: """ Returns the duration of the video in seconds. @@ -504,16 +592,28 @@ class VideoFromFile(VideoInput): metadata: Optional[dict] = None, bit_depth: int | None = None, crf: float | None = None, + color_space: str | None = None, ): + if color_space is not None and color_space not in VIDEO_COLOR_TRANSFERS: + raise ValueError(f"Unsupported video color space: {color_space}") + _, output_format, _ = video_output_config(path, format, codec) if isinstance(self.__file, io.BytesIO): self.__file.seek(0) # Reset the BytesIO object to the beginning with av.open(self.__file, mode='r') as container: container_format = container.format.name video_stream = container.streams.video[0] if len(container.streams.video) > 0 else None - video_encoding = video_stream.codec.name if video_stream is not None else None + video_encoding = video_stream.codec.canonical_name if video_stream is not None else None source_bit_depth = video_stream_bit_depth(video_stream) + source_color_space = video_stream_color_space(video_stream) + if source_color_space is not None and color_space is not None and source_color_space != color_space: + raise ValueError( + f"Cannot save {source_color_space} video as {color_space} without color conversion; " + f"use auto or {source_color_space}" + ) reuse_streams = True - if format != VideoContainer.AUTO and format not in container_format.split(","): + if format != VideoContainer.AUTO and VIDEO_CONTAINER_FORMATS[VideoContainer(format)] not in container_format.split(","): + reuse_streams = False + if output_format == VideoContainer.WEBM and not webm_streams_compatible(container.streams): reuse_streams = False if codec != VideoCodec.AUTO and codec != video_encoding and video_encoding is not None: reuse_streams = False @@ -521,6 +621,8 @@ class VideoFromFile(VideoInput): reuse_streams = False if crf is not None: reuse_streams = False + if color_space is not None: + reuse_streams = False if self.__start_time or self.__duration: reuse_streams = False if self.__crop is not None: @@ -529,7 +631,7 @@ class VideoFromFile(VideoInput): if not reuse_streams: if bit_depth is None: bit_depth = source_bit_depth - return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth, crf=crf) + return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth, crf=crf, color_space=color_space) streams = container.streams @@ -563,9 +665,10 @@ class VideoFromFile(VideoInput): metadata: dict | None, bit_depth: int, crf: float | None = None, + color_space: str | None = None, ): - """Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length.""" - open_kwargs = mp4_output_open_kwargs(path, format, codec) + """Re-encode one frame at a time; peak memory does not scale with video length.""" + open_kwargs, output_format, output_codec = video_output_config(path, format, codec) video_stream = self._get_first_video_stream(container) start_time, duration = self.get_active_trim_window() start_pts = int(start_time / video_stream.time_base) @@ -580,6 +683,10 @@ class VideoFromFile(VideoInput): container.seek(start_pts, stream=video_stream) audio_stream = last_decodable_audio_stream(container) + source_color_space = video_stream_color_space(video_stream) + preserve_source_color = source_color_space is not None + if color_space in HDR_COLOR_TRANSFERS or source_color_space in HDR_COLOR_TRANSFERS: + bit_depth = max(bit_depth, 10) pix_fmt = "yuv420p10le" if bit_depth >= 10 else "yuv420p" rate = Fraction(video_stream.average_rate) if video_stream.average_rate else Fraction(1) @@ -599,6 +706,8 @@ class VideoFromFile(VideoInput): logging.warning("Audio stream parameters could not be determined; ignoring audio.") audio_stream = None if audio_stream is not None: + if output_format == VideoContainer.WEBM: + sample_rate = 48000 audio_time_base = Fraction(1, sample_rate) layout = {1: "mono", 2: "stereo", 6: "5.1"}.get(channels, "stereo") resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=sample_rate) @@ -708,24 +817,28 @@ class VideoFromFile(VideoInput): if crop_rect is not None: out_width, out_height = crop_rect[2], crop_rect[3] if out_width % 2 or out_height % 2: - raise ValueError(f"H.264 output requires even dimensions, got {out_width}x{out_height}") + raise ValueError(f"{output_codec.value.upper()} output requires even dimensions, got {out_width}x{out_height}") source_size = (frame.width, frame.height) output = av.open(path, **open_kwargs) # Add metadata before writing any streams write_output_metadata(container, output, metadata) - out_video = output.add_stream("h264", rate=rate) + out_video = output.add_stream(VIDEO_ENCODERS[output_codec], rate=rate) # no B-frames: reordering makes mp4 sample durations follow decode order, # so irregular-VFR spans and trim windows land wrong out_video.codec_context.max_b_frames = 0 out_video.width = out_width out_video.height = out_height out_video.pix_fmt = pix_fmt - if crf is not None: - out_video.options = {"crf": str(crf)} + out_video.options = video_encoder_options(output_codec, crf) + if preserve_source_color: + copy_color_properties(video_stream, out_video.codec_context) + elif color_space is not None: + set_video_color_properties(out_video.codec_context, color_space) # source pts pass through (rebased to 0), so variable frame rate survives out_video.codec_context.time_base = video_stream.time_base if audio_stream is not None: - out_audio = output.add_stream("aac", rate=sample_rate, layout=layout) + audio_codec = "libopus" if output_format == VideoContainer.WEBM else "aac" + out_audio = output.add_stream(audio_codec, rate=sample_rate, layout=layout) if (frame.width, frame.height) != source_size: # encoding would silently rescale the new geometry into the old one raise ValueError( @@ -763,11 +876,15 @@ class VideoFromFile(VideoInput): crop_filter = (g, g_src, g_sink) crop_filter[1].push(frame) frame = crop_filter[2].pull() - if frame.color_range == ColorRange.JPEG: + if frame.color_range == ColorRange.JPEG and not preserve_source_color: # compress full-range sources (yuvj/MJPEG) to limited range frame = frame.reformat(format=pix_fmt, src_color_range="JPEG", dst_color_range="MPEG") else: frame = frame.reformat(format=pix_fmt) + if preserve_source_color: + copy_color_properties(video_stream, frame) + elif color_space is not None: + set_video_color_properties(frame, color_space) frame_output_end = None if frame.pts is not None: if video_pts_offset is None: @@ -930,6 +1047,9 @@ class VideoFromComponents(VideoInput): def get_bit_depth(self) -> int: return self.__bit_depth + def get_color_space(self) -> str: + return "sRGB" + def save_to( self, path: str, @@ -938,12 +1058,19 @@ class VideoFromComponents(VideoInput): metadata: Optional[dict] = None, bit_depth: int | None = None, crf: float | None = None, + color_space: str | None = None, ): """Save the video to a file path or BytesIO buffer.""" - open_kwargs = mp4_output_open_kwargs(path, format, codec) + if color_space is None: + color_space = "sRGB" + if color_space is not None and color_space not in VIDEO_COLOR_TRANSFERS: + raise ValueError(f"Unsupported video color space: {color_space}") + open_kwargs, output_format, output_codec = video_output_config(path, format, codec) # None means "use the depth this video was created with" (CreateVideo's choice). if bit_depth is None: bit_depth = self.__bit_depth + if color_space in HDR_COLOR_TRANSFERS: + bit_depth = max(bit_depth, 10) is_10bit = bit_depth >= 10 with av.open(path, **open_kwargs) as output: # Add metadata before writing any streams @@ -954,22 +1081,28 @@ class VideoFromComponents(VideoInput): frame_rate = Fraction(round(self.__components.frame_rate * 1000), 1000) # Create a video stream pix_fmt = "yuv420p10le" if is_10bit else "yuv420p" - video_stream = output.add_stream('h264', rate=frame_rate) + video_stream = output.add_stream(VIDEO_ENCODERS[output_codec], rate=frame_rate) video_stream.width = self.__components.images.shape[2] video_stream.height = self.__components.images.shape[1] video_stream.pix_fmt = pix_fmt - if crf is not None: - video_stream.options = {"crf": str(crf)} + video_stream.options = video_encoder_options(output_codec, crf) + if color_space is not None: + set_video_color_properties(video_stream.codec_context, color_space) # Create an audio stream audio_sample_rate = 1 + audio_resampler = None audio_stream: Optional[av.AudioStream] = None if self.__components.audio: - audio_sample_rate = int(self.__components.audio['sample_rate']) + source_audio_sample_rate = int(self.__components.audio['sample_rate']) + audio_sample_rate = 48000 if output_format == VideoContainer.WEBM else source_audio_sample_rate waveform = self.__components.audio['waveform'] - waveform = waveform[0, :, :math.ceil((audio_sample_rate / frame_rate) * self.__components.images.shape[0])] + waveform = waveform[0, :, :math.ceil((source_audio_sample_rate / frame_rate) * self.__components.images.shape[0])] layout = {1: 'mono', 2: 'stereo', 6: '5.1'}.get(waveform.shape[0], 'stereo') - audio_stream = output.add_stream('aac', rate=audio_sample_rate, layout=layout) + audio_codec = "libopus" if output_format == VideoContainer.WEBM else "aac" + audio_stream = output.add_stream(audio_codec, rate=audio_sample_rate, layout=layout) + if audio_sample_rate != source_audio_sample_rate: + audio_resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=audio_sample_rate) # Encode video for i, frame in enumerate(self.__components.images): @@ -980,7 +1113,14 @@ class VideoFromComponents(VideoInput): else: img = (frame * 255).clamp(0, 255).byte().cpu().numpy() # shape: (H, W, 3) frame = av.VideoFrame.from_ndarray(img, format='rgb24') - frame = frame.reformat(format=pix_fmt) + dst_colorspace = None + if color_space == "sRGB": + dst_colorspace = BT709_NCL + elif color_space in HDR_COLOR_TRANSFERS: + dst_colorspace = BT2020_NCL + frame = frame.reformat(format=pix_fmt, dst_colorspace=dst_colorspace) + if color_space is not None: + set_video_color_properties(frame, color_space) packet = video_stream.encode(frame) output.mux(packet) @@ -990,9 +1130,14 @@ class VideoFromComponents(VideoInput): if audio_stream and self.__components.audio: frame = av.AudioFrame.from_ndarray(waveform.float().cpu().contiguous().numpy(), format='fltp', layout=layout) - frame.sample_rate = audio_sample_rate + frame.sample_rate = source_audio_sample_rate frame.pts = 0 - output.mux(audio_stream.encode(frame)) + frames = [frame] if audio_resampler is None else audio_resampler.resample(frame) + for frame in frames: + output.mux(audio_stream.encode(frame)) + if audio_resampler is not None: + for frame in audio_resampler.resample(None): + output.mux(audio_stream.encode(frame)) # Flush encoder output.mux(audio_stream.encode(None)) diff --git a/comfy_api/latest/_util/video_types.py b/comfy_api/latest/_util/video_types.py index 2f8ae887d..faa190c33 100644 --- a/comfy_api/latest/_util/video_types.py +++ b/comfy_api/latest/_util/video_types.py @@ -7,6 +7,7 @@ from .._input import ImageInput, AudioInput, MaskInput class VideoCodec(str, Enum): AUTO = "auto" H264 = "h264" + AV1 = "av1" @classmethod def as_input(cls) -> list[str]: @@ -18,6 +19,8 @@ class VideoCodec(str, Enum): class VideoContainer(str, Enum): AUTO = "auto" MP4 = "mp4" + MKV = "mkv" + WEBM = "webm" @classmethod def as_input(cls) -> list[str]: @@ -35,6 +38,10 @@ class VideoContainer(str, Enum): value = cls(value) if value == VideoContainer.MP4 or value == VideoContainer.AUTO: return "mp4" + if value == VideoContainer.MKV: + return "mkv" + if value == VideoContainer.WEBM: + return "webm" return "" @dataclass diff --git a/comfy_api_nodes/apis/bfl.py b/comfy_api_nodes/apis/bfl.py index 389706cf4..0e33f2f5e 100644 --- a/comfy_api_nodes/apis/bfl.py +++ b/comfy_api_nodes/apis/bfl.py @@ -166,3 +166,13 @@ class Flux3VideoContinuationRequest(Flux3VideoRequest): start_video: str = Field( ..., description="MP4 (URL or base64); the new clip carries on from its final frames." ) + + +class BFLFluxVideoUpscaleRequest(BaseModel): + model_config = ConfigDict(extra="forbid") + + input_video: str = Field(..., description="MP4 (URL or base64), 1 to 20 seconds.") + upscale_factor: float = Field(2.0, ge=1.5, le=3.0) + creativity: int = Field(1, description="0 preserves the source precisely, 1 enhances detail.") + prompt: str | None = Field(None) + safety_tolerance: int = Field(2, ge=0, le=4) diff --git a/comfy_api_nodes/apis/bria.py b/comfy_api_nodes/apis/bria.py index 7a98428c3..f55de74bc 100644 --- a/comfy_api_nodes/apis/bria.py +++ b/comfy_api_nodes/apis/bria.py @@ -57,6 +57,81 @@ class BriaRemoveBackgroundRequest(BaseModel): seed: int = Field(...) +class BriaGenFillRequest(BaseModel): + image: str = Field(...) + mask: str = Field( + ..., + description="Binary mask defining the region to fill: white (255) pixels are generated, " + "black (0) pixels are preserved. Must have the same aspect ratio as the image.", + ) + prompt: str = Field(...) + negative_prompt: str | None = Field(None) + refine_prompt: bool = Field(True) + seed: int = Field(...) + prompt_content_moderation: bool = Field(False, description="If true, returns 422 on prompt moderation failure.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image or mask moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaEraseRequest(BaseModel): + image: str = Field(...) + mask: str = Field( + ..., + description="Binary mask defining the region to erase: white (255) pixels are removed, " + "black (0) pixels are preserved. Must have the same aspect ratio as the image.", + ) + mask_type: str = Field("manual", description="'manual' for hand-drawn masks, 'automatic' for segmentation masks.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image or mask moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaExpandRequest(BaseModel): + image: str = Field(...) + aspect_ratio: str | float | None = Field( + None, + description="Target ratio: a preset string (1:1, 2:3, 3:2, 3:4, 4:3, 4:5, 5:4, 9:16, 16:9) " + "or a float between 0.5 and 3.0. When set, the canvas/placement fields are ignored.", + ) + canvas_size: list[int] | None = Field(None, description="Output canvas [width, height]; area up to 5000x5000.") + original_image_size: list[int] | None = Field( + None, description="Size [width, height] of the original image inside the canvas." + ) + original_image_location: list[int] | None = Field( + None, + description="Top-left corner [x, y] of the original image inside the canvas; " + "values may fall outside the canvas, cropping the image.", + ) + prompt: str | None = Field(None, description="If omitted, Bria auto-generates a prompt from the image.") + negative_prompt: str | None = Field(None) + seed: int = Field(...) + prompt_content_moderation: bool = Field(False, description="If true, returns 422 on prompt moderation failure.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + +class BriaIncreaseResolutionRequest(BaseModel): + image: str = Field(...) + desired_increase: int = Field(..., description="Resolution multiplier, 2 or 4.") + visual_input_content_moderation: bool = Field( + False, description="If true, returns 422 on image moderation failure." + ) + visual_output_content_moderation: bool = Field( + False, description="If true, returns 422 on visual output moderation failure." + ) + + class BriaStatusResponse(BaseModel): request_id: str = Field(...) status_url: str = Field(...) @@ -72,6 +147,26 @@ class BriaRemoveBackgroundResponse(BaseModel): result: BriaRemoveBackgroundResult | None = Field(None) +class BriaImageResult(BaseModel): + image_url: str = Field(...) + + +class BriaImageResultResponse(BaseModel): + status: str = Field(...) + result: BriaImageResult | None = Field(None) + + +class BriaExpandResult(BaseModel): + image_url: str = Field(...) + prompt: str | None = Field(None) + seed: int | None = Field(None) + + +class BriaExpandResponse(BaseModel): + status: str = Field(...) + result: BriaExpandResult | None = Field(None) + + class BriaImageEditResult(BaseModel): structured_prompt: str = Field(...) image_url: str = Field(...) diff --git a/comfy_api_nodes/apis/bytedance.py b/comfy_api_nodes/apis/bytedance.py index 7ee83e5f3..87d22034a 100644 --- a/comfy_api_nodes/apis/bytedance.py +++ b/comfy_api_nodes/apis/bytedance.py @@ -18,7 +18,8 @@ class Seedream4Options(BaseModel): class Seedream5OptimizePromptOptions(BaseModel): - thinking: Literal["auto", "enabled", "disabled"] = Field(...) + thinking: Literal["auto", "enabled", "disabled"] | None = Field(None) + mode: Literal["standard", "fast"] | None = Field(None) class Seedream4TaskCreationRequest(BaseModel): @@ -116,6 +117,7 @@ class Seedance2TaskCreationRequest(BaseModel): seed: int | None = Field(None, ge=0, le=2147483647) watermark: bool | None = Field(None) output_format: str | None = Field(None) + omni_reference_task_type: str | None = Field(None, description="One of: auto, reference, edit, extend.") class TaskCreationResponse(BaseModel): @@ -186,35 +188,6 @@ class SeedanceVirtualLibraryCreateAssetRequest(BaseModel): asset_type: str | None = Field(None, description="BytePlus asset type. Defaults to Image server-side when omitted.") -# Dollars per 1K tokens, keyed by (model_id, has_video_input, resolution). -SEEDANCE2_PRICE_PER_1K_TOKENS = { - ("dreamina-seedance-2-0-260128", False, "480p"): 0.007, - ("dreamina-seedance-2-0-260128", True, "480p"): 0.0043, - ("dreamina-seedance-2-0-260128", False, "720p"): 0.007, - ("dreamina-seedance-2-0-260128", True, "720p"): 0.0043, - ("dreamina-seedance-2-0-260128", False, "1080p"): 0.0077, - ("dreamina-seedance-2-0-260128", True, "1080p"): 0.0047, - ("dreamina-seedance-2-0-260128", False, "4k"): 0.004, - ("dreamina-seedance-2-0-260128", True, "4k"): 0.0024, - ("dreamina-seedance-2-0-fast-260128", False, "480p"): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True, "480p"): 0.0033, - ("dreamina-seedance-2-0-fast-260128", False, "720p"): 0.0056, - ("dreamina-seedance-2-0-fast-260128", True, "720p"): 0.0033, - ("dreamina-seedance-2-0-mini", False, "480p"): 0.0035, - ("dreamina-seedance-2-0-mini", True, "480p"): 0.0021, - ("dreamina-seedance-2-0-mini", False, "720p"): 0.0035, - ("dreamina-seedance-2-0-mini", True, "720p"): 0.0021, - ("dreamina-seedance-2-5-260628", False, "480p"): 0.0107, - ("dreamina-seedance-2-5-260628", True, "480p"): 0.0064, - ("dreamina-seedance-2-5-260628", False, "720p"): 0.0107, - ("dreamina-seedance-2-5-260628", True, "720p"): 0.0064, -} - - -def seedance2_price_per_1k_tokens(model_id: str, has_video_input: bool, resolution: str) -> float | None: - return SEEDANCE2_PRICE_PER_1K_TOKENS.get((model_id, has_video_input, resolution)) - - RECOMMENDED_PRESETS = [ ("1024x1024 (1:1)", 1024, 1024), ("864x1152 (3:4)", 864, 1152), @@ -329,6 +302,7 @@ SEEDANCE2_REF_VIDEO_PIXEL_LIMITS = { "dreamina-seedance-2-5-260628": { "480p": {"min": 409_600, "max": 8_295_044}, "720p": {"min": 409_600, "max": 8_295_044}, + "1080p": {"min": 409_600, "max": 8_295_044}, }, } diff --git a/comfy_api_nodes/apis/fishaudio.py b/comfy_api_nodes/apis/fishaudio.py new file mode 100644 index 000000000..f615d3d24 --- /dev/null +++ b/comfy_api_nodes/apis/fishaudio.py @@ -0,0 +1,49 @@ +from pydantic import BaseModel, Field + + +class FishAudioProsody(BaseModel): + speed: float = Field(1.0, description="Speaking rate multiplier, 0.5-2.0") + volume: float = Field(0.0, description="Volume adjustment in decibels") + + +class FishAudioTTSRequest(BaseModel): + text: str = Field(..., description="Text to synthesize") + reference_id: str | list[str] | None = Field(None, description="Voice model ID or list of IDs") + temperature: float = Field(0.7, description="Expressiveness, 0-1") + top_p: float = Field(0.7, description="Nucleus sampling diversity, (0, 1]") + prosody: FishAudioProsody = Field(..., description="Speed and volume adjustments") + normalize: bool = Field(True, description="Normalize numbers and text for English and Chinese") + format: str = Field("wav", description="Output audio format") + + +class FishAudioASRRequest(BaseModel): + language: str | None = Field(None, description="Optional ISO 639-1 language hint") + ignore_timestamps: bool = Field(True, description="Skip precise timestamp computation") + + +class FishAudioASRSegment(BaseModel): + text: str | None = Field(None, description="Segment text") + start: float | None = Field(None, description="Segment start time in seconds") + end: float | None = Field(None, description="Segment end time in seconds") + + +class FishAudioASRResponse(BaseModel): + text: str | None = Field(None, description="Transcribed text") + duration: float | None = Field(None, description="Audio duration in seconds") + segments: list[FishAudioASRSegment] | None = Field(None, description="Timestamped transcript segments") + language_code: str | None = Field(None, description="Detected language as ISO 639-1 code") + language: str | None = Field(None, description="Detected language display name") + + +class FishAudioCreateModelRequest(BaseModel): + type: str = Field("tts", description="Model type") + title: str = Field(..., description="Voice model name") + train_mode: str = Field("fast", description="Training mode; fast is instantly available") + visibility: str = Field("private", description="Model visibility") + enhance_audio_quality: bool = Field(..., description="Enhance reference audio quality") + + +class FishAudioCreateModelResponse(BaseModel): + id: str = Field(..., alias="_id", description="Voice model ID for use as reference_id") + state: str | None = Field(None, description="Training state") + visibility: str | None = Field(None, description="Model visibility") diff --git a/comfy_api_nodes/apis/grok.py b/comfy_api_nodes/apis/grok.py index 526d8c8ab..82dfc8b82 100644 --- a/comfy_api_nodes/apis/grok.py +++ b/comfy_api_nodes/apis/grok.py @@ -9,6 +9,7 @@ class ImageGenerationRequest(BaseModel): seed: int = Field(...) response_format: str = Field("url") resolution: str = Field(...) + quality: str | None = Field(None) class InputUrlObject(BaseModel): @@ -28,6 +29,7 @@ class ImageEditRequest(BaseModel): seed: int = Field(...) response_format: str = Field("url") aspect_ratio: str | None = Field(...) + quality: str | None = Field(None) class VideoGenerationRequest(BaseModel): diff --git a/comfy_api_nodes/apis/minimax.py b/comfy_api_nodes/apis/minimax.py index bac4572d4..12a0853ac 100644 --- a/comfy_api_nodes/apis/minimax.py +++ b/comfy_api_nodes/apis/minimax.py @@ -161,12 +161,30 @@ class Hailuo03TaskCreationRequest(BaseModel): ..., min_length=1 ) resolution: str = Field(...) - duration: int = Field(..., ge=5, le=15) + duration: int = Field(..., ge=4, le=15) ratio: str | None = Field(None) seed: int | None = Field(None, ge=0, le=4294967295) aigc_watermark: bool | None = Field(None) +class Hailuo03ContextIRRequest(BaseModel): + model: str = Field(...) + content: list[Hailuo03TextContent | Hailuo03ImageContent | Hailuo03VideoContent | Hailuo03AudioContent] = Field( + ..., min_length=1 + ) + duration: int = Field(..., ge=4, le=15) + ratio: str | None = Field(None) + + +class Hailuo03RegenerationRequest(BaseModel): + model: str = Field(...) + content: list[Hailuo03TextContent | Hailuo03ImageContent | Hailuo03VideoContent | Hailuo03AudioContent] = Field( + ..., min_length=1 + ) + resolution: str = Field(...) + aigc_watermark: bool | None = Field(None) + + class Hailuo03TaskCreationResponse(BaseModel): task_id: str = Field(...) @@ -178,6 +196,7 @@ class Hailuo03TaskError(BaseModel): class Hailuo03TaskContent(BaseModel): url: str | None = Field(None) + prompt: str | None = Field(None) class Hailuo03TaskUsage(BaseModel): diff --git a/comfy_api_nodes/apis/qwen.py b/comfy_api_nodes/apis/qwen.py new file mode 100644 index 000000000..90b68dee8 --- /dev/null +++ b/comfy_api_nodes/apis/qwen.py @@ -0,0 +1,46 @@ +from pydantic import BaseModel, Field + + +class QwenImageContentItem(BaseModel): + image: str | None = Field(None) + text: str | None = Field(None) + + +class QwenImageMessage(BaseModel): + role: str = Field("user") + content: list[QwenImageContentItem] = Field(...) + + +class QwenImageInputField(BaseModel): + messages: list[QwenImageMessage] = Field(...) + + +class QwenImageParametersField(BaseModel): + size: str | None = Field(None, description="Output resolution as 'width*height'; omit for the model default.") + n: int = Field(1, ge=1, le=6) + seed: int = Field(..., ge=0, le=2147483647) + prompt_extend: bool = Field(True) + watermark: bool = Field(False) + negative_prompt: str | None = Field(None) + + +class QwenImageGenerationRequest(BaseModel): + model: str = Field(...) + input: QwenImageInputField = Field(...) + parameters: QwenImageParametersField = Field(...) + + +class QwenImageChoice(BaseModel): + finish_reason: str | None = Field(None) + message: QwenImageMessage | None = Field(None) + + +class QwenImageOutputField(BaseModel): + choices: list[QwenImageChoice] = Field(default_factory=list) + + +class QwenImageGenerationResponse(BaseModel): + output: QwenImageOutputField | None = Field(None) + request_id: str = Field(...) + code: str | None = Field(None, description="Error code for the failed request.") + message: str | None = Field(None, description="Details about the failed request.") diff --git a/comfy_api_nodes/nodes_bfl.py b/comfy_api_nodes/nodes_bfl.py index 6c961dc0c..dd8653cdd 100644 --- a/comfy_api_nodes/nodes_bfl.py +++ b/comfy_api_nodes/nodes_bfl.py @@ -13,18 +13,19 @@ from comfy_api_nodes.apis.bfl import ( BFLFluxProGenerateResponse, BFLFluxProUltraGenerateRequest, BFLFluxStatusResponse, + BFLFluxVideoUpscaleRequest, BFLFluxVTORequest, BFLStatus, Flux2ProGenerateRequest, Flux3ImageToVideoRequest, Flux3TextToVideoRequest, Flux3VideoContinuationRequest, - Flux3VideoRequest, ) from comfy_api_nodes.util import ( ApiEndpoint, convert_mask_to_image, download_url_to_image_tensor, + downscale_video_to_max_pixels, download_url_to_video_output, get_number_of_images, poll_op, @@ -36,6 +37,8 @@ from comfy_api_nodes.util import ( validate_aspect_ratio_string, validate_image_dimensions, validate_string, + validate_video_dimensions, + validate_video_duration, ) @@ -588,16 +591,12 @@ class FluxEraseNode(IO.ComfyNode): ), ) - def price_extractor(_r: BaseModel) -> float | None: - return None if initial_response.cost is None else initial_response.cost / 100 - response = await poll_op( cls, ApiEndpoint(initial_response.polling_url), response_model=BFLFluxStatusResponse, status_extractor=lambda r: r.status, progress_extractor=lambda r: r.progress, - price_extractor=price_extractor, completed_statuses=[BFLStatus.ready], failed_statuses=[ BFLStatus.request_moderated, @@ -669,16 +668,12 @@ class FluxVTONode(IO.ComfyNode): ), ) - def price_extractor(_r: BaseModel) -> float | None: - return None if initial_response.cost is None else initial_response.cost / 100 - response = await poll_op( cls, ApiEndpoint(initial_response.polling_url), response_model=BFLFluxStatusResponse, status_extractor=lambda r: r.status, progress_extractor=lambda r: r.progress, - price_extractor=price_extractor, completed_statuses=[BFLStatus.ready], failed_statuses=[ BFLStatus.request_moderated, @@ -802,16 +797,12 @@ class Flux2ProImageNode(IO.ComfyNode): ), ) - def price_extractor(_r: BaseModel) -> float | None: - return None if initial_response.cost is None else initial_response.cost / 100 - response = await poll_op( cls, ApiEndpoint(initial_response.polling_url), response_model=BFLFluxStatusResponse, status_extractor=lambda r: r.status, progress_extractor=lambda r: r.progress, - price_extractor=price_extractor, completed_statuses=[BFLStatus.ready], failed_statuses=[ BFLStatus.request_moderated, @@ -994,16 +985,12 @@ class Flux2ImageNode(IO.ComfyNode): ), ) - def price_extractor(_r: BaseModel) -> float | None: - return None if initial_response.cost is None else initial_response.cost / 100 - response = await poll_op( cls, ApiEndpoint(initial_response.polling_url), response_model=BFLFluxStatusResponse, status_extractor=lambda r: r.status, progress_extractor=lambda r: r.progress, - price_extractor=price_extractor, completed_statuses=[BFLStatus.ready], failed_statuses=[ BFLStatus.request_moderated, @@ -1164,24 +1151,26 @@ class Flux3VideoNodeBase(IO.ComfyNode): ) -async def _flux3_execute(cls: type[IO.ComfyNode], request: Flux3VideoRequest) -> IO.NodeOutput: - initial_response = await sync_op( - cls, - ApiEndpoint(path="/proxy/bfl/v1/flux-3-video", method="POST"), - response_model=BFLFluxProGenerateResponse, - data=request, +_FLUX3_VIDEO_ENDPOINT = ApiEndpoint(path="/proxy/bfl/v1/flux-3-video", method="POST") +_FLUX_VIDEO_UPSCALE_ENDPOINT = ApiEndpoint(path="/proxy/bfl/v1/flux-tools/video-upscale-v1", method="POST") +_BFL_POLL_PROXY_PATH = "/proxy/bfl/v1/get_result" + + +async def _bfl_video_execute( + cls: type[IO.ComfyNode], endpoint: ApiEndpoint, request: BaseModel, poll_via_proxy: bool = False +) -> IO.NodeOutput: + initial_response = await sync_op(cls, endpoint, response_model=BFLFluxProGenerateResponse, data=request) + poll_endpoint = ( + ApiEndpoint(path=_BFL_POLL_PROXY_PATH, query_params={"polling_url": initial_response.polling_url}) + if poll_via_proxy + else ApiEndpoint(initial_response.polling_url) ) - - def price_extractor(_r: BaseModel) -> float | None: - return None if initial_response.cost is None else initial_response.cost / 100 - response = await poll_op( cls, - ApiEndpoint(initial_response.polling_url), + poll_endpoint, response_model=BFLFluxStatusResponse, status_extractor=lambda r: r.status, progress_extractor=lambda r: r.progress, - price_extractor=price_extractor, completed_statuses=[BFLStatus.ready], failed_statuses=[ BFLStatus.request_moderated, @@ -1243,7 +1232,7 @@ class Flux3TextToVideoNode(Flux3VideoNodeBase): request = Flux3TextToVideoRequest( **cls.common_fields(prompt, aspect_ratio, duration, resolution, generate_audio, safety_tolerance) ) - return await _flux3_execute(cls, request) + return await _bfl_video_execute(cls, _FLUX3_VIDEO_ENDPOINT, request) class Flux3ImageToVideoNode(Flux3VideoNodeBase): @@ -1341,7 +1330,7 @@ class Flux3ImageToVideoNode(Flux3VideoNodeBase): keyframes=list(zip(times, urls)) if times is not None else urls, **fields, ) - return await _flux3_execute(cls, request) + return await _bfl_video_execute(cls, _FLUX3_VIDEO_ENDPOINT, request) class Flux3VideoContinuationNode(Flux3VideoNodeBase): @@ -1392,7 +1381,131 @@ class Flux3VideoContinuationNode(Flux3VideoNodeBase): fields = cls.common_fields(prompt, aspect_ratio, duration, resolution, generate_audio, safety_tolerance) url = await upload_video_to_comfyapi(cls, video, wait_label="Uploading source video") request = Flux3VideoContinuationRequest(start_video=url, **fields) - return await _flux3_execute(cls, request) + return await _bfl_video_execute(cls, _FLUX3_VIDEO_ENDPOINT, request) + + +_FLUX_VIDEO_UPSCALE_MODES = {"creative": 1, "precise": 0} +_FLUX_VIDEO_UPSCALE_MAX_INPUT_PIXELS = 3840 * 2160 +_FLUX_VIDEO_UPSCALE_MAX_ASPECT_RATIO = 4.0 + + +class FluxVideoUpscaleNode(IO.ComfyNode): + + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FluxVideoUpscaleNode", + display_name="Flux Video Upscale", + category="partner/video/BFL", + description="Upscales a video 1.5 to 3 times with FLUX super-resolution, either preserving " + "the source precisely or creatively enhancing its detail.", + inputs=[ + IO.Video.Input( + "video", + tooltip="Source clip of 1 to 20 seconds with an aspect ratio between 1:4 and 4:1. " + "The output is rendered at 24 fps and capped at about 14.4 megapixels per frame.", + ), + IO.Float.Input( + "upscale_factor", + default=2.0, + min=1.5, + max=3.0, + step=0.1, + tooltip="Output size relative to the source. Very large sources are upscaled by " + "less than the requested factor because of the per-frame cap.", + ), + IO.Combo.Input( + "mode", + options=list(_FLUX_VIDEO_UPSCALE_MODES), + default="creative", + tooltip="'creative' restores and invents fine detail, best for generated footage, " + "textures and scenery. 'precise' sharpens the source without changing it, " + "for faces, products and real footage.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Optional description of the clip that steers the enhanced detail. " + "Leave empty for a neutral upscale.", + ), + IO.Boolean.Input( + "auto_downscale", + default=True, + tooltip="Automatically downscale sources larger than 3840x2160 pixels in area to fit " + "the input limit. Aspect ratio is preserved; smaller videos are untouched.", + ), + IO.Int.Input( + "safety_tolerance", + default=2, + min=0, + max=4, + advanced=True, + tooltip="Moderation tolerance, 0 is the strictest.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=0xFFFFFFFF, + control_after_generate=True, + tooltip="Seed to determine if node should re-run; FLUX picks its own seed, so " + "actual results are nondeterministic regardless of this value.", + ), + ], + outputs=[IO.Video.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["mode"]), + expr=""" + ( + $precise := widgets.mode = "precise"; + {"type":"range_usd", + "min_usd": $precise ? 0.212 : 0.297, + "max_usd": $precise ? 0.848 : 1.188, + "format": {"approximate": true, "suffix": "/s", "note": "(1080p-4K output)"}} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + video: Input.Video, + upscale_factor: float, + mode: str, + prompt: str, + auto_downscale: bool, + safety_tolerance: int, + seed: int, + ) -> IO.NodeOutput: + validate_video_duration(video, min_duration=1.0, max_duration=20.0) + validate_video_dimensions(video, min_width=64, min_height=64) + width, height = video.get_dimensions() + if max(width, height) > _FLUX_VIDEO_UPSCALE_MAX_ASPECT_RATIO * min(width, height): + raise ValueError(f"Video aspect ratio must be between 1:4 and 4:1, got {width}x{height}.") + if auto_downscale: + video = downscale_video_to_max_pixels(video, _FLUX_VIDEO_UPSCALE_MAX_INPUT_PIXELS) + elif width * height > _FLUX_VIDEO_UPSCALE_MAX_INPUT_PIXELS: + raise ValueError( + f"Video must be at most 3840x2160 pixels in area, got {width}x{height}. " + "Enable auto_downscale or use a smaller video." + ) + url = await upload_video_to_comfyapi(cls, video, wait_label="Uploading source video") + request = BFLFluxVideoUpscaleRequest( + input_video=url, + upscale_factor=round(upscale_factor, 1), + creativity=_FLUX_VIDEO_UPSCALE_MODES[mode], + prompt=prompt.strip() or None, + safety_tolerance=safety_tolerance, + ) + return await _bfl_video_execute(cls, _FLUX_VIDEO_UPSCALE_ENDPOINT, request, poll_via_proxy=True) class BFLExtension(ComfyExtension): @@ -1412,6 +1525,7 @@ class BFLExtension(ComfyExtension): Flux3TextToVideoNode, Flux3ImageToVideoNode, Flux3VideoContinuationNode, + FluxVideoUpscaleNode, ] diff --git a/comfy_api_nodes/nodes_bria.py b/comfy_api_nodes/nodes_bria.py index 77f780a3b..90cade2d0 100644 --- a/comfy_api_nodes/nodes_bria.py +++ b/comfy_api_nodes/nodes_bria.py @@ -6,7 +6,13 @@ from typing_extensions import override from comfy_api.latest import IO, ComfyExtension, Input from comfy_api_nodes.apis.bria import ( BriaEditImageRequest, + BriaEraseRequest, + BriaExpandRequest, + BriaExpandResponse, + BriaGenFillRequest, BriaImageEditResponse, + BriaImageResultResponse, + BriaIncreaseResolutionRequest, BriaRemoveBackgroundRequest, BriaRemoveBackgroundResponse, BriaRemoveVideoBackgroundRequest, @@ -21,13 +27,30 @@ from comfy_api_nodes.util import ( convert_mask_to_image, download_url_to_image_tensor, download_url_to_video_output, + downscale_image_tensor_by_max_side, + get_image_dimensions, poll_op, sync_op, upload_image_to_comfyapi, upload_video_to_comfyapi, + validate_string, validate_video_duration, ) +BRIA_MAX_OUTPUT_SIDE = 8192 +BRIA_MIN_RATIO = 0.5 +BRIA_MAX_RATIO = 3.0 +BRIA_MIN_SHORT_SIDE = 224 + + +def _upscaled_output_side(height: int, width: int, multiplier: int) -> int: + prescale = max(1.0, BRIA_MIN_SHORT_SIDE / min(height, width)) + return round(max(height, width) * prescale * multiplier) + + +def _smallest_output_side(height: int, width: int, multiplier: int) -> int: + return round(max(height, width) / min(height, width) * BRIA_MIN_SHORT_SIDE * multiplier) + class BriaImageEditNode(IO.ComfyNode): @@ -243,6 +266,503 @@ class BriaRemoveImageBackground(IO.ComfyNode): return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) +def _mask_to_binary_image(mask: Input.Image, action: str) -> torch.Tensor: + binary = (mask > 0.5).float() + if not binary.any(): + raise ValueError( + f"The mask is empty, so there is nothing to {action}. Masks are binarized at 50%: " + f"areas painted at less than half opacity are ignored." + ) + return convert_mask_to_image(binary) + + +def _validate_mask_aspect_ratio(image: Input.Image, mask: Input.Image) -> None: + ih, iw = image.shape[1], image.shape[2] + mh, mw = mask.shape[-2], mask.shape[-1] + if abs(iw * mh - ih * mw) > 0.01 * ih * mw: + raise ValueError(f"Mask must have the same aspect ratio as the image: image is {iw}x{ih}, mask is {mw}x{mh}.") + + +class BriaGenFill(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaGenFill", + display_name="Bria Generative Fill", + category="partner/image/Bria", + description="Generate objects or scenery inside a masked region of an image using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.Mask.Input( + "mask", + tooltip="White areas are filled with generated content, black areas are preserved. " + "The mask is binarized before sending, so partially painted areas count as white. " + "Must have the same aspect ratio as the image.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Description of what to generate inside the masked region.", + ), + IO.String.Input("negative_prompt", multiline=True, default=""), + IO.Boolean.Input( + "refine_prompt", + default=True, + tooltip="Automatically adjust the prompt for better results; " + "disable to use the prompt exactly as written.", + ), + IO.Int.Input( + "seed", + default=42, + min=1, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("prompt_content_moderation", default=False), + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0429}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + mask: Input.Image, + prompt: str, + negative_prompt: str, + refine_prompt: bool, + seed: int, + moderation: InputModerationSettings, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1) + _validate_mask_aspect_ratio(image, mask) + mask_image = _mask_to_binary_image(mask, "fill") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/gen_fill", method="POST"), + data=BriaGenFillRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + mask=await upload_image_to_comfyapi( + cls, mask_image, total_pixels=None, wait_label="Uploading mask" + ), + prompt=prompt, + negative_prompt=negative_prompt if negative_prompt else None, + refine_prompt=refine_prompt, + seed=seed, + prompt_content_moderation=moderation.get("prompt_content_moderation", False), + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + +class BriaEraser(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaEraser", + display_name="Bria Eraser", + category="partner/image/Bria", + description="Remove objects or areas outlined by a mask from an image using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.Mask.Input( + "mask", + tooltip="White areas are erased, black areas are preserved. " + "The mask is binarized before sending, so partially painted areas count as white. " + "Must have the same aspect ratio as the image.", + ), + IO.Combo.Input( + "mask_type", + options=["manual", "automatic"], + tooltip="manual for hand-drawn or brush masks, " + "automatic for masks produced by segmentation models such as SAM.", + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + mask: Input.Image, + mask_type: str, + moderation: dict, + ) -> IO.NodeOutput: + _validate_mask_aspect_ratio(image, mask) + mask_image = _mask_to_binary_image(mask, "erase") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/erase", method="POST"), + data=BriaEraseRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + mask=await upload_image_to_comfyapi( + cls, mask_image, total_pixels=None, wait_label="Uploading mask" + ), + mask_type=mask_type, + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + +class BriaExpandImage(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaExpandImage", + display_name="Bria Expand Image", + category="partner/image/Bria", + description="Expand an image beyond its borders with generated content using Bria.", + inputs=[ + IO.Image.Input("image"), + IO.DynamicCombo.Input( + "expand_mode", + options=[ + *[IO.DynamicCombo.Option(ratio, []) for ratio in + ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"]], + IO.DynamicCombo.Option( + "custom_ratio", + [ + IO.Int.Input( + "ratio_width", + default=21, + min=1, + max=100, + tooltip="Width side of the target ratio: 21 and 9 give 21:9.", + ), + IO.Int.Input( + "ratio_height", + default=9, + min=1, + max=100, + tooltip="Height side of the target ratio: 21 and 9 give 21:9. " + f"Bria only accepts width/height between {BRIA_MIN_RATIO} and " + f"{BRIA_MAX_RATIO}, so anything taller than 1:2 needs the manual mode.", + ), + ], + ), + IO.DynamicCombo.Option( + "manual", + [ + IO.Int.Input("canvas_width", default=1000, min=64, max=5000), + IO.Int.Input("canvas_height", default=1000, min=64, max=5000), + IO.Int.Input( + "image_width", + default=500, + min=1, + max=5000, + tooltip="Width of the original image inside the canvas.", + ), + IO.Int.Input( + "image_height", + default=500, + min=1, + max=5000, + tooltip="Height of the original image inside the canvas.", + ), + IO.Int.Input( + "image_x", + default=250, + min=-5000, + max=5000, + tooltip="X position of the image's top-left corner inside the canvas; " + "may fall outside the canvas, cropping the image.", + ), + IO.Int.Input( + "image_y", + default=250, + min=-5000, + max=5000, + tooltip="Y position of the image's top-left corner inside the canvas; " + "may fall outside the canvas, cropping the image.", + ), + ], + ), + ], + tooltip="Target shape of the expanded image: a preset aspect ratio, a custom ratio, " + "or manual placement of the original image on a canvas. " + "Manual is the only mode that can reach a canvas taller than 1:2.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Optional description of the expanded scene; " + "when empty, Bria generates one from the image.", + ), + IO.String.Input("negative_prompt", multiline=True, default=""), + IO.Int.Input( + "seed", + default=42, + min=1, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("prompt_content_moderation", default=False), + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[ + IO.Image.Output(), + IO.String.Output(display_name="prompt", tooltip="The prompt used for the expansion; " + "auto-generated by Bria when the prompt input is empty."), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + expand_mode: dict, + prompt: str, + negative_prompt: str, + seed: int, + moderation: InputModerationSettings, + ) -> IO.NodeOutput: + mode = expand_mode["expand_mode"] + aspect_ratio = canvas_size = original_image_size = original_image_location = None + if mode == "manual": + canvas_size = [expand_mode["canvas_width"], expand_mode["canvas_height"]] + original_image_size = [expand_mode["image_width"], expand_mode["image_height"]] + original_image_location = [expand_mode["image_x"], expand_mode["image_y"]] + elif mode == "custom_ratio": + ratio_width, ratio_height = expand_mode["ratio_width"], expand_mode["ratio_height"] + aspect_ratio = ratio_width / ratio_height + if not BRIA_MIN_RATIO <= aspect_ratio <= BRIA_MAX_RATIO: + raise ValueError( + f"Bria accepts a width-to-height ratio between {BRIA_MIN_RATIO} and {BRIA_MAX_RATIO}: " + f"{ratio_width}:{ratio_height} is {aspect_ratio:.4f}. " + f"Use the manual expand mode to reach a canvas of any shape." + ) + else: + aspect_ratio = mode + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/expand", method="POST"), + data=BriaExpandRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + aspect_ratio=aspect_ratio, + canvas_size=canvas_size, + original_image_size=original_image_size, + original_image_location=original_image_location, + prompt=prompt if prompt else None, + negative_prompt=negative_prompt if negative_prompt else None, + seed=seed, + prompt_content_moderation=moderation.get("prompt_content_moderation", False), + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaExpandResponse, + ) + return IO.NodeOutput( + await download_url_to_image_tensor(response.result.image_url), + response.result.prompt or "", + ) + + +class BriaIncreaseResolution(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="BriaIncreaseResolution", + display_name="Bria Increase Resolution", + category="partner/image/Bria", + description="Upscale an image by 2x or 4x using Bria, preserving the original content.", + inputs=[ + IO.Image.Input("image"), + IO.Combo.Input( + "desired_increase", + options=["2", "4"], + tooltip="Resolution multiplier. The output must fit within 8192 pixels on each side.", + ), + IO.Boolean.Input( + "auto_downscale", + default=False, + tooltip="Automatically lower the multiplier, and downscale the input image if that is " + "still not enough, when the output would exceed the limit.", + ), + IO.DynamicCombo.Input( + "moderation", + options=[ + IO.DynamicCombo.Option("false", []), + IO.DynamicCombo.Option( + "true", + [ + IO.Boolean.Input("visual_input_moderation", default=False), + IO.Boolean.Input("visual_output_moderation", default=False), + ], + ), + ], + tooltip="Moderation settings", + ), + ], + outputs=[IO.Image.Output()], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.0286}""", + ), + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + desired_increase: str, + auto_downscale: bool, + moderation: dict, + ) -> IO.NodeOutput: + multiplier = int(desired_increase) + height, width = get_image_dimensions(image) + if _upscaled_output_side(height, width, multiplier) > BRIA_MAX_OUTPUT_SIDE: + candidates = [c for c in (4, 2) if c <= multiplier] + if not auto_downscale: + predicted = _upscaled_output_side(height, width, multiplier) + raise ValueError( + f"Bria can upscale up to a maximum output dimension of {BRIA_MAX_OUTPUT_SIDE} pixels: " + f"input is {width}x{height}, x{multiplier} would be {predicted} pixels on the long side. " + f"Enable auto_downscale, or use a smaller input image or a lower multiplier." + ) + fitted = next( + (c for c in candidates if _upscaled_output_side(height, width, c) <= BRIA_MAX_OUTPUT_SIDE), None + ) + if fitted is not None: + multiplier = fitted + else: + shrinkable = next((c for c in sorted(candidates) if _smallest_output_side(height, width, c) + <= BRIA_MAX_OUTPUT_SIDE), None) + if shrinkable is None: + raise ValueError( + f"This image cannot be upscaled by Bria at any multiplier: it is {width}x{height}, and " + f"Bria first enlarges the short side to {BRIA_MIN_SHORT_SIDE} pixels, which pushes the " + f"long side past the {BRIA_MAX_OUTPUT_SIDE} pixel limit. Crop it to a squarer shape first." + ) + multiplier = shrinkable + image = downscale_image_tensor_by_max_side(image, max_side=BRIA_MAX_OUTPUT_SIDE // multiplier) + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/bria/v2/image/edit/increase_resolution", method="POST"), + data=BriaIncreaseResolutionRequest( + image=await upload_image_to_comfyapi(cls, image, total_pixels=None, wait_label="Uploading image"), + desired_increase=multiplier, + visual_input_content_moderation=moderation.get("visual_input_moderation", False), + visual_output_content_moderation=moderation.get("visual_output_moderation", False), + ), + response_model=BriaStatusResponse, + ) + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/bria/v2/status/{response.request_id}"), + status_extractor=lambda r: r.status, + response_model=BriaImageResultResponse, + ) + return IO.NodeOutput(await download_url_to_image_tensor(response.result.image_url)) + + class BriaRemoveVideoBackground(IO.ComfyNode): @classmethod @@ -572,6 +1092,10 @@ class BriaExtension(ComfyExtension): return [ BriaImageEditNode, BriaRemoveImageBackground, + BriaGenFill, + BriaEraser, + BriaExpandImage, + BriaIncreaseResolution, BriaRemoveVideoBackground, BriaVideoGreenScreen, BriaVideoReplaceBackground, diff --git a/comfy_api_nodes/nodes_bytedance.py b/comfy_api_nodes/nodes_bytedance.py index 09fc445d8..1f2dfd21b 100644 --- a/comfy_api_nodes/nodes_bytedance.py +++ b/comfy_api_nodes/nodes_bytedance.py @@ -49,13 +49,14 @@ from comfy_api_nodes.apis.bytedance import ( TaskVideoContentUrl, Text2ImageTaskCreationRequest, Text2VideoTaskCreationRequest, - seedance2_price_per_1k_tokens, seedance2_reference_limits, ) from comfy_api_nodes.util import ( ApiEndpoint, audio_bytes_to_audio_input, audio_input_to_mp3, + bytesio_to_image_tensor, + download_url_as_bytesio, download_url_to_image_tensor, download_url_to_video_output, downscale_image_tensor_by_max_side, @@ -116,7 +117,7 @@ SEEDANCE_MODELS = { SEEDANCE_MODEL_TOOLTIP = ( "Seedance 2.5 for the newest model, videos up to 30 seconds and mp4/mov output; " - "Seedance 2.0 for maximum quality and 1080p/4k; Fast for speed optimization; " + "Seedance 2.0 for maximum quality and 4k; Fast for speed optimization; " "Mini for the fastest, lowest-cost generation." ) @@ -406,20 +407,6 @@ async def _seedance_virtual_library_upload_video_asset( return f"asset://{create_resp.asset_id}" -def _seedance2_price_extractor(model_id: str, has_video_input: bool, resolution: str): - """Returns a price_extractor closure for Seedance 2.0 poll_op.""" - rate = seedance2_price_per_1k_tokens(model_id, has_video_input, resolution) - if rate is None: - return None - - def extractor(response: TaskStatusResponse) -> float | None: - if response.usage is None: - return None - return response.usage.total_tokens * 1.43 * rate / 1_000.0 - - return extractor - - def get_image_url_from_response(response: ImageTaskCreationResponse) -> str: if response.error: error_msg = f"ByteDance request failed. Code: {response.error['code']}, message: {response.error['message']}" @@ -768,6 +755,8 @@ def _seedream_model_inputs( max_width: int = 6240, max_height: int = 4992, supports_batch: bool = True, + supports_fast: bool = False, + include_common: bool = False, ): inputs = [ IO.Combo.Input( @@ -828,16 +817,282 @@ def _seedream_model_inputs( advanced=True, ) ) + if supports_fast: + inputs.append( + IO.Combo.Input( + "prompt_optimization", + options=["standard", "fast"], + default="standard", + tooltip="Prompt-optimization mode when reference images are provided: " + "'standard' gives higher quality, 'fast' shorter generation time.", + advanced=True, + ) + ) + if include_common: + inputs.extend( + [ + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "watermark", + default=False, + 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)." + ), + advanced=True, + ), + ] + ) return inputs -class ByteDanceSeedreamNodeV2(IO.ComfyNode): +class ByteDanceSeedreamNodeV3(IO.ComfyNode): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="ByteDanceSeedreamNodeV3", + display_name="ByteDance Seedream 4.5 & 5.0", + category="partner/image/ByteDance", + description="Unified text-to-image generation and precise single-sentence editing at up to 4K resolution.", + inputs=[ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Text prompt for creating or editing an image.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "seedream 5.0 pro", + _seedream_model_inputs( + max_ref_images=10, + presets=RECOMMENDED_PRESETS_SEEDREAM_5_PRO, + max_width=3136, + max_height=2496, + supports_batch=False, + supports_fast=True, + include_common=True, + ), + ), + IO.DynamicCombo.Option( + "seedream 5.0 lite", + _seedream_model_inputs( + max_ref_images=14, + presets=RECOMMENDED_PRESETS_SEEDREAM_5_LITE, + include_common=True, + ), + ), + IO.DynamicCombo.Option( + "seedream-4-5-251128", + _seedream_model_inputs( + max_ref_images=10, + presets=RECOMMENDED_PRESETS_SEEDREAM_4_5, + include_common=True, + ), + ), + IO.DynamicCombo.Option( + "seedream-4-0-250828", + _seedream_model_inputs( + max_ref_images=10, + presets=RECOMMENDED_PRESETS_SEEDREAM_4_0, + include_common=True, + ), + ), + ], + ), + ], + outputs=[ + IO.Image.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + widgets=["model", "model.size_preset", "model.width", "model.height"], + input_groups=["model.images"], + ), + expr=""" + ( + $model := $string(widgets.model); + $sp := $string($lookup(widgets, "model.size_preset")); + $w := $lookup(widgets, "model.width"); + $h := $lookup(widgets, "model.height"); + $px := ($type($w) = "number" and $type($h) = "number") ? $w * $h : 0; + $refs := $lookup(inputGroups, "model.images"); + $extra := ($type($refs) = "number" and $refs > 1) ? ($refs - 1) * 0.003 : 0; + $isPro := $contains($model, "5.0 pro"); + $isCustom := $contains($sp, "custom"); + $sizeKnown := $isCustom ? $px > 0 : ($contains($sp, "1k") or $contains($sp, "2k")); + $proPrice := $isCustom + ? ($px < 2610000 ? 0.045 : 0.09) + : ($contains($sp, "1k") ? 0.045 : 0.09); + ($isPro and ($sizeKnown = false)) + ? { + "type": "range_usd", + "min_usd": 0.045 + $extra, + "max_usd": 0.09 + $extra, + "format": { "suffix": "/Image", "approximate": true } + } + : { + "type": "usd", + "usd": $isPro ? $proPrice + $extra + : $contains($model, "5.0 lite") ? 0.035 + : $contains($model, "4-5") ? 0.04 + : 0.03, + "format": { "suffix": $isPro ? "/Image" : " x images/Run", "approximate": true } + } + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + prompt: str, + 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"]] + presets = SEEDREAM_PRESETS[model_id] + is_pro = "seedream-5-0-pro" in model_id + + size_preset = model.get("size_preset", presets[0][0]) + width = model.get("width", 2048) + height = model.get("height", 2048) + max_images = model.get("max_images", 1) + sequential_image_generation = "disabled" if max_images == 1 else "auto" + images_dict = model.get("images") or {} + fail_on_partial = model.get("fail_on_partial", False) + prompt_optimization = model.get("prompt_optimization", "standard") + seed = model.get("seed", seed) + watermark = model.get("watermark", watermark) + thinking = model.get("thinking", thinking) + + w = h = None + for label, tw, th in presets: + if label == size_preset: + w, h = tw, th + break + if w is None or h is None: + w, h = width, height + + out_num_pixels = w * h + mp_provided = out_num_pixels / 1_000_000.0 + if is_pro: + if out_num_pixels < 921_600: + raise ValueError( + f"Minimum image resolution for the selected model is 0.92MP, but {mp_provided:.2f}MP provided." + ) + if out_num_pixels > 4_194_304: + raise ValueError( + f"Maximum image resolution for the selected model is 4.19MP, but {mp_provided:.2f}MP provided." + ) + else: + if ("seedream-4-5" in model_id or "seedream-5-0" in model_id) and out_num_pixels < 3_686_400: + raise ValueError( + f"Minimum image resolution for the selected model is 3.68MP, but {mp_provided:.2f}MP provided." + ) + if "seedream-4-0" in model_id and out_num_pixels < 921_600: + raise ValueError( + f"Minimum image resolution that the selected model can generate is 0.92MP, " + f"but {mp_provided:.2f}MP provided." + ) + if out_num_pixels > 16_777_216: + raise ValueError( + f"Maximum image resolution for the selected model is 16.78MP, but {mp_provided:.2f}MP provided." + ) + + image_tensors: list[Input.Image] = [t for t in images_dict.values() if t is not None] + n_input_images = sum(get_number_of_images(t) for t in image_tensors) + max_num_of_images = 14 if model_id == "seedream-5-0-260128" else 10 + if n_input_images > max_num_of_images: + raise ValueError( + f"Maximum of {max_num_of_images} reference images are supported, but {n_input_images} received." + ) + if sequential_image_generation == "auto" and n_input_images + max_images > 15: + 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: + for tensor in image_tensors: + validate_image_aspect_ratio(tensor, (1, 3), (3, 1)) + reference_images_urls = await upload_images_to_comfyapi( + cls, + image_tensors, + max_images=n_input_images, + mime_type="image/png", + wait_label="Uploading reference images", + ) + + optimize_prompt_options = None + if n_input_images == 0: + optimize_prompt_options = Seedream5OptimizePromptOptions(thinking="enabled" if thinking else "disabled") + elif prompt_optimization == "fast": + optimize_prompt_options = Seedream5OptimizePromptOptions(mode="fast") + response = await sync_op( + cls, + ApiEndpoint(path=BYTEPLUS_IMAGE_ENDPOINT, method="POST"), + response_model=ImageTaskCreationResponse, + data=Seedream4TaskCreationRequest( + model=model_id, + prompt=prompt, + image=reference_images_urls, + size=f"{w}x{h}", + seed=seed, + 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: + return IO.NodeOutput(await download_url_to_image_tensor(get_image_url_from_response(response))) + urls = [str(d["url"]) for d in response.data if isinstance(d, dict) and "url" in d] + if fail_on_partial and len(urls) < len(response.data): + raise RuntimeError(f"Only {len(urls)} of {len(response.data)} images were generated before error.") + return IO.NodeOutput(torch.cat([await download_url_to_image_tensor(i) for i in urls])) + + +class ByteDanceSeedreamNodeV2(ByteDanceSeedreamNodeV3): @classmethod def define_schema(cls): return IO.Schema( node_id="ByteDanceSeedreamNodeV2", - display_name="ByteDance Seedream 4.5 & 5.0", + display_name="ByteDance Seedream 4.5 & 5.0 (Legacy)", category="partner/image/ByteDance", description="Unified text-to-image generation and precise single-sentence editing at up to 4K resolution.", inputs=[ @@ -911,6 +1166,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): IO.Hidden.unique_id, ], is_api_node=True, + is_deprecated=True, price_badge=IO.PriceBadge( depends_on=IO.PriceBadgeDepends( widgets=["model", "model.size_preset", "model.width", "model.height"] @@ -939,117 +1195,6 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode): ), ) - @classmethod - async def execute( - cls, - prompt: str, - 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"]] - presets = SEEDREAM_PRESETS[model_id] - is_pro = "seedream-5-0-pro" in model_id - - size_preset = model.get("size_preset", presets[0][0]) - width = model.get("width", 2048) - height = model.get("height", 2048) - max_images = model.get("max_images", 1) - sequential_image_generation = "disabled" if max_images == 1 else "auto" - images_dict = model.get("images") or {} - fail_on_partial = model.get("fail_on_partial", False) - - w = h = None - for label, tw, th in presets: - if label == size_preset: - w, h = tw, th - break - if w is None or h is None: - w, h = width, height - - out_num_pixels = w * h - mp_provided = out_num_pixels / 1_000_000.0 - if is_pro: - if out_num_pixels < 921_600: - raise ValueError( - f"Minimum image resolution for the selected model is 0.92MP, but {mp_provided:.2f}MP provided." - ) - if out_num_pixels > 4_194_304: - raise ValueError( - f"Maximum image resolution for the selected model is 4.19MP, but {mp_provided:.2f}MP provided." - ) - else: - if ("seedream-4-5" in model_id or "seedream-5-0" in model_id) and out_num_pixels < 3_686_400: - raise ValueError( - f"Minimum image resolution for the selected model is 3.68MP, but {mp_provided:.2f}MP provided." - ) - if "seedream-4-0" in model_id and out_num_pixels < 921_600: - raise ValueError( - f"Minimum image resolution that the selected model can generate is 0.92MP, " - f"but {mp_provided:.2f}MP provided." - ) - if out_num_pixels > 16_777_216: - raise ValueError( - f"Maximum image resolution for the selected model is 16.78MP, but {mp_provided:.2f}MP provided." - ) - - image_tensors: list[Input.Image] = [t for t in images_dict.values() if t is not None] - n_input_images = sum(get_number_of_images(t) for t in image_tensors) - max_num_of_images = 14 if model_id == "seedream-5-0-260128" else 10 - if n_input_images > max_num_of_images: - raise ValueError( - f"Maximum of {max_num_of_images} reference images are supported, but {n_input_images} received." - ) - if sequential_image_generation == "auto" and n_input_images + max_images > 15: - 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: - for tensor in image_tensors: - validate_image_aspect_ratio(tensor, (1, 3), (3, 1)) - reference_images_urls = await upload_images_to_comfyapi( - cls, - image_tensors, - max_images=n_input_images, - mime_type="image/png", - 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"), - response_model=ImageTaskCreationResponse, - data=Seedream4TaskCreationRequest( - model=model_id, - prompt=prompt, - image=reference_images_urls, - size=f"{w}x{h}", - seed=seed, - 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: - return IO.NodeOutput(await download_url_to_image_tensor(get_image_url_from_response(response))) - urls = [str(d["url"]) for d in response.data if isinstance(d, dict) and "url" in d] - if fail_on_partial and len(urls) < len(response.data): - raise RuntimeError(f"Only {len(urls)} of {len(response.data)} images were generated before error.") - return IO.NodeOutput(torch.cat([await download_url_to_image_tensor(i) for i in urls])) - - class ByteDanceSeedreamLayerSeparationNode(IO.ComfyNode): @classmethod @@ -1315,7 +1460,9 @@ class ByteDanceSeedreamLayerSeparationNode(IO.ComfyNode): left, top, rect_w, rect_h = spec["left"], spec["top"], spec["rect_w"], spec["rect_h"] async with semaphore: try: - rgba = (await download_url_to_image_tensor(str(item["url"])))[0] + # the layer math below needs the alpha channel, and ByteDance encodes + # alpha-less images as plain RGB (the base plate is one), so force RGBA + rgba = bytesio_to_image_tensor(await download_url_as_bytesio(str(item["url"])), mode="RGBA")[0] except ProcessingInterrupted: raise except Exception as exc: @@ -2069,7 +2216,7 @@ def _seedance2_text_inputs(resolutions: list[str], default_ratio: str = "16:9"): ] -def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = False): +def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = False, with_task_type: bool = False): return [ IO.String.Input( "prompt", @@ -2080,7 +2227,7 @@ def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = ), IO.Combo.Input( "resolution", - options=["480p", "720p"], + options=["480p", "720p", "1080p"], default="720p", tooltip="Resolution of the output video.", ), @@ -2124,6 +2271,29 @@ def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = if with_video_editing else [] ), + *( + [ + IO.Combo.Input( + "task_type", + options=["auto", "reference", "edit", "extend"], + default="auto", + tooltip="What to do with the reference media. Every value except auto is " + "validated when the task is submitted, so mismatched settings fail before " + "generation starts. auto: the model infers the task from the prompt and " + "inputs, and settings that conflict with its reading fail only after " + "generation has started. reference: generate a new video guided by the " + "reference images, videos, and audio. edit: change a connected reference " + "video (add, remove, replace); the output keeps the source clip's own length " + "and aspect ratio, and the duration and ratio widgets are ignored. extend: " + "continue a connected reference video forward or backward; the prompt should " + "say 'extend forward', 'extend backward', or 'continue', the aspect ratio " + "follows the source clip, and the output contains only the newly generated " + "segment of the duration you set, not the source clip.", + ) + ] + if with_task_type + else [] + ), IO.Combo.Input( "output_format", options=["mp4"], @@ -2133,9 +2303,9 @@ def _seedance25_text_inputs(with_ratio: bool = True, with_video_editing: bool = ] -def _seedance25_reference_inputs(): +def _seedance25_reference_inputs(with_video_editing: bool = False, with_task_type: bool = False): return [ - *_seedance25_text_inputs(with_video_editing=True), + *_seedance25_text_inputs(with_video_editing=with_video_editing, with_task_type=with_task_type), IO.Autogrow.Input( "reference_images", template=IO.Autogrow.TemplateNames( @@ -2196,17 +2366,23 @@ def _seedance2_build_request( watermark: bool, ratio: str, ) -> Seedance2TaskCreationRequest: - video_editing = bool(model.get("video_editing")) + task_type = model.get("task_type", "auto") + duration = model["duration"] + if model.get("video_editing") or task_type == "edit": + ratio, duration = "adaptive", -1 + elif task_type == "extend": + ratio = "adaptive" return Seedance2TaskCreationRequest( model=model_id, content=content, generate_audio=model["generate_audio"], resolution=model["resolution"], - ratio="adaptive" if video_editing else ratio, - duration=-1 if video_editing else model["duration"], + ratio=ratio, + duration=duration, seed=seed, watermark=watermark, output_format=model.get("output_format"), + omni_reference_task_type=None if task_type == "auto" else task_type, ) @@ -2216,18 +2392,21 @@ _SEEDANCE2_PRICE_EXPR_TEMPLATE = """ $res := $lookup(widgets, "model.resolution"); $ratio := $lookup(widgets, "model.ratio"); $dur := $lookup(widgets, "model.duration"); - $auto := $lookup(widgets, "model.video_editing") = true; + $auto := __IS_EDIT__; $hasVideo := __HAS_VIDEO__; $ready := $type($m) = "string" and $type($res) = "string" and ($auto or $type($dur) = "number"); $ready ? ( $contains($m, "2.5") ? ( $is480 := $res = "480p"; - $perFrame := $ratio = "1:1" ? ($is480 ? 400 : 900) : - $ratio = "4:3" ? ($is480 ? 411.25 : 905.6719) : - $ratio = "3:4" ? ($is480 ? 411.25 : 905.6719) : - $ratio = "21:9" ? ($is480 ? 418.5 : 904.3945) : - ($is480 ? 400.3125 : 900); - $price := $hasVideo ? 0.009152 : 0.015301; + $is1080 := $res = "1080p"; + $perFrame := $ratio = "1:1" ? ($is480 ? 400 : $is1080 ? 2025 : 900) : + $ratio = "4:3" ? ($is480 ? 411.25 : $is1080 ? 2028 : 905.6719) : + $ratio = "3:4" ? ($is480 ? 411.25 : $is1080 ? 2028 : 905.6719) : + $ratio = "21:9" ? ($is480 ? 418.5 : $is1080 ? 2037.9648 : 904.3945) : + ($is480 ? 400.3125 : $is1080 ? 2025 : 900); + $price := $is1080 + ? ($hasVideo ? 0.01001 : 0.016731) + : ($hasVideo ? 0.009152 : 0.015301); $costFor := function($d) { $floor($perFrame * (24 * $d + 1)) / 1000 * $price }; $lo := $costFor($auto ? 4 : $dur); $hi := $costFor(($auto ? 30 : $dur) + ($hasVideo ? 30 : 0)); @@ -2261,14 +2440,13 @@ _SEEDANCE2_PRICE_EXPR_TEMPLATE = """ _SEEDANCE_AUDIO_POLICY_CODE = "OutputAudioSensitiveContentDetected.PolicyViolation" _SEEDANCE_TASK_TYPE_CONSTRAINT_CODE = "InvalidParameter.TaskTypeConstraint" +_SEEDANCE_TASK_TYPE_MISMATCH_CODE = "InvalidParameter.TaskTypeMismatch" async def _seedance2_poll_video_task( cls: type[IO.ComfyNode], task_id: str, - model_id: str, - resolution: str, - has_video_input: bool, + task_type: str | None = None, ) -> TaskStatusResponse: try: return await poll_op( @@ -2276,9 +2454,6 @@ async def _seedance2_poll_video_task( ApiEndpoint(path=f"{BYTEPLUS_SEEDANCE2_TASK_STATUS_ENDPOINT}/{task_id}"), response_model=TaskStatusResponse, status_extractor=lambda r: r.status, - price_extractor=_seedance2_price_extractor( - model_id, has_video_input=has_video_input, resolution=resolution - ), poll_interval=9, ) except Exception as exc: @@ -2289,19 +2464,48 @@ async def _seedance2_poll_video_task( "to get a silent video, or adjust the prompt and try again." ) from exc if _SEEDANCE_TASK_TYPE_CONSTRAINT_CODE in str(exc): + if task_type is None: + raise ValueError( + "Seedance read this prompt as editing the reference video, and an edit always " + "takes its duration and aspect ratio from that video. Enable video_editing on " + "this node and run again, or reword the prompt so it describes a new video " + "rather than a change to the reference one." + ) from exc + if task_type == "edit": + raise ValueError( + "The request does not satisfy the 'edit' constraints: the clip being edited " + "must be 4 to 30 seconds long." + ) from exc + if task_type == "extend": + raise ValueError( + "The request does not satisfy the 'extend' constraints: the clip being " + "extended must be 1.9 to 30 seconds long." + ) from exc raise ValueError( - "Seedance read this prompt as editing the reference video, and an edit always " - "takes its duration and aspect ratio from that video. Enable video_editing on " - "this node and run again, or reword the prompt so it describes a new video " - "rather than a change to the reference one." + "Seedance decided from the prompt that this task's duration or aspect ratio " + "must come from the reference video, and the current settings conflict with " + "that. Set task_type to the task you mean ('edit' or 'extend') and run again, " + "or reword the prompt so it describes a new video rather than a change to the " + "reference one." + ) from exc + if _SEEDANCE_TASK_TYPE_MISMATCH_CODE in str(exc): + raise ValueError( + f"Seedance read this prompt as a different task than the selected task_type " + f"'{task_type}'. Reword the prompt so it matches: an extend prompt should say " + "'extend forward', 'extend backward', or 'continue'; an edit prompt should use " + "words like add, remove, replace, or change. Or set task_type to auto." ) from exc raise -def _seedance2_price_badge(with_reference_videos: bool) -> IO.PriceBadge: +def _seedance2_price_badge(with_reference_videos: bool, legacy_video_editing: bool = False) -> IO.PriceBadge: widgets = ["model", "model.resolution", "model.ratio", "model.duration"] + if legacy_video_editing: + is_edit = '$lookup(widgets, "model.video_editing") = true' + else: + is_edit = '$lookup(widgets, "model.task_type") = "edit"' if with_reference_videos: - widgets.append("model.video_editing") + widgets.append("model.video_editing" if legacy_video_editing else "model.task_type") has_video = ( '$exists(inputGroups) and $lookup(inputGroups, "model.reference_videos") > 0' if with_reference_videos @@ -2312,7 +2516,7 @@ def _seedance2_price_badge(with_reference_videos: bool) -> IO.PriceBadge: widgets=widgets, input_groups=["model.reference_videos"] if with_reference_videos else [], ), - expr=_SEEDANCE2_PRICE_EXPR_TEMPLATE.replace("__HAS_VIDEO__", has_video), + expr=_SEEDANCE2_PRICE_EXPR_TEMPLATE.replace("__HAS_VIDEO__", has_video).replace("__IS_EDIT__", is_edit), ) @@ -2388,9 +2592,7 @@ class ByteDance2TextToVideoNode(IO.ComfyNode): ), response_model=TaskCreationResponse, ) - response = await _seedance2_poll_video_task( - cls, initial_response.id, model_id, model["resolution"], has_video_input=False - ) + response = await _seedance2_poll_video_task(cls, initial_response.id) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -2581,9 +2783,7 @@ class ByteDance2FirstLastFrameNode(IO.ComfyNode): data=_seedance2_build_request(model, model_id, content, seed, watermark, ratio=request_ratio), response_model=TaskCreationResponse, ) - response = await _seedance2_poll_video_task( - cls, initial_response.id, model_id, model["resolution"], has_video_input=False - ) + response = await _seedance2_poll_video_task(cls, initial_response.id) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) @@ -2662,12 +2862,12 @@ def _seedance2_reference_inputs(resolutions: list[str], default_ratio: str = "16 ] -class ByteDance2ReferenceNode(IO.ComfyNode): +class ByteDance2ReferenceNodeV2(IO.ComfyNode): @classmethod def define_schema(cls): return IO.Schema( - node_id="ByteDance2ReferenceNode", + node_id="ByteDance2ReferenceNodeV2", display_name="ByteDance Seedance 2.5 Reference to Video", category="partner/video/ByteDance", description="Generate, edit, or extend video using Seedance 2.5 or 2.0 with reference " @@ -2676,7 +2876,7 @@ class ByteDance2ReferenceNode(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ - IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs()), + IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs(with_task_type=True)), IO.DynamicCombo.Option( "Seedance 2.0", _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), @@ -2761,6 +2961,13 @@ class ByteDance2ReferenceNode(IO.ComfyNode): f"(videos={len(reference_videos)}, video assets={len(reference_video_assets)}). " f"Maximum is {limits['max_videos']}." ) + task_type = model.get("task_type") + if task_type in ("edit", "extend") and total_videos == 0: + raise ValueError( + f"A '{task_type}' task needs at least one reference video. Connect the video " + f"you want to {'change' if task_type == 'edit' else 'continue'}, or set " + "task_type to 'reference' to generate a new video from the references you have." + ) total_audios = len(reference_audios) + len(reference_audio_assets) if total_audios > limits["max_audios"]: raise ValueError( @@ -2772,8 +2979,6 @@ class ByteDance2ReferenceNode(IO.ComfyNode): for key in reference_images: reference_images[key] = _prepare_seedance_image(reference_images[key]) - has_video_input = total_videos > 0 - if model.get("auto_downscale") and reference_videos: max_px = SEEDANCE2_REF_VIDEO_PIXEL_LIMITS.get(model_id, {}).get(model["resolution"], {}).get("max") if max_px: @@ -2893,11 +3098,75 @@ class ByteDance2ReferenceNode(IO.ComfyNode): response_model=TaskCreationResponse, ) response = await _seedance2_poll_video_task( - cls, initial_response.id, model_id, model["resolution"], has_video_input=has_video_input + cls, + initial_response.id, + task_type=task_type, ) return IO.NodeOutput(await download_url_to_video_output(response.content.video_url)) +class ByteDance2ReferenceNode(ByteDance2ReferenceNodeV2): + + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="ByteDance2ReferenceNode", + display_name="ByteDance Seedance 2.5 Reference to Video (Legacy)", + category="partner/video/ByteDance", + description="Generate, edit, or extend video using Seedance 2.5 or 2.0 with reference " + "images, videos, and audio. Supports multimodal reference, video editing, and video extension.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option("Seedance 2.5", _seedance25_reference_inputs(with_video_editing=True)), + IO.DynamicCombo.Option( + "Seedance 2.0", + _seedance2_reference_inputs(["480p", "720p", "1080p", "4k"], default_ratio="adaptive"), + ), + IO.DynamicCombo.Option( + "Seedance 2.0 Fast", + _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), + ), + IO.DynamicCombo.Option( + "Seedance 2.0 Mini", + _seedance2_reference_inputs(["480p", "720p"], default_ratio="adaptive"), + ), + ], + tooltip=SEEDANCE_MODEL_TOOLTIP, + ), + IO.Int.Input( + "seed", + default=0, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add a watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + is_deprecated=True, + price_badge=_seedance2_price_badge(with_reference_videos=True, legacy_video_editing=True), + ) + + async def process_video_task( cls: type[IO.ComfyNode], payload: Text2VideoTaskCreationRequest | Image2VideoTaskCreationRequest, @@ -3405,6 +3674,7 @@ class ByteDanceExtension(ComfyExtension): ByteDanceImageNode, ByteDanceSeedreamNode, ByteDanceSeedreamNodeV2, + ByteDanceSeedreamNodeV3, ByteDanceSeedreamLayerSeparationNode, ByteDanceTextToVideoNode, ByteDanceImageToVideoNode, @@ -3413,6 +3683,7 @@ class ByteDanceExtension(ComfyExtension): ByteDance2TextToVideoNode, ByteDance2FirstLastFrameNode, ByteDance2ReferenceNode, + ByteDance2ReferenceNodeV2, ByteDanceCreateImageAsset, ByteDanceCreateVideoAsset, ByteDanceSeedAudioNode, diff --git a/comfy_api_nodes/nodes_fishaudio.py b/comfy_api_nodes/nodes_fishaudio.py new file mode 100644 index 000000000..d40fc2740 --- /dev/null +++ b/comfy_api_nodes/nodes_fishaudio.py @@ -0,0 +1,454 @@ +import json +import re +import uuid + +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension, Input +from comfy_api_nodes.apis.fishaudio import ( + FishAudioASRRequest, + FishAudioASRResponse, + FishAudioCreateModelRequest, + FishAudioCreateModelResponse, + FishAudioProsody, + FishAudioTTSRequest, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + audio_bytes_to_audio_input, + audio_ndarray_to_bytesio, + audio_tensor_to_contiguous_ndarray, + sync_op, + sync_op_raw, + validate_string, +) + +FISHAUDIO_VOICE = "FISHAUDIO_VOICE" + +FISHAUDIO_VOICES = [ + ("802e3bc2b27e49c2995d23ef70e6ac89", "Energetic Male (en)"), + ("b545c585f631496c914815291da4e893", "Friendly Women (en)"), + ("933563129e564b19a115bedd57b7406a", "Sarah (en)"), + ("8d21b053e2804e2a890e1cf62f267b6f", "Verity (en)"), + ("f48d143a59a946ab87c0130fd081f349", "Polo (en)"), + ("bf322df2096a46f18c579d0baa36f41d", "Adrian (en)"), + ("98655a12fa944e26b274c535e5e03842", "E-girl (en)"), + ("0327fdb5da9e4fd782899a8058c8ae2b", "Narrator (en)"), + ("5212eb29e500460391d03af42af6552e", "Warm Conversational Voice (en)"), + ("5c8dc6a69c0b4edfb32634db6384bf34", "Warm Storyteller (en)"), + ("7a18a1851d2649108c48ec9f2c80eb2c", "Dramatic Character Male (en)"), + ("59cb5986671546eaa6ca8ae6f29f6d22", "News Narrator (zh)"), + ("bf6c479f5a384b8d857310030035824b", "Lively Female (zh)"), + ("faccba1a8ac54016bcfc02761285e67f", "Gentle Female (zh)"), + ("5161d41404314212af1254556477c17d", "Energetic Female (ja)"), + ("0089dce5fefb4c6ba9b9f2f0debe1ddc", "Calm Female (ja)"), + ("45c5d3723c9c42f598e4776dcfd5f02d", "Calm Male (ja)"), +] + +FISHAUDIO_VOICE_MAP = {label: voice_id for voice_id, label in FISHAUDIO_VOICES} + +MAX_REFERENCE_AUDIO_SECONDS = 270 + + +def _rewrite_voice_tags(text: str, voice_count: int) -> tuple[str, set[int]]: + referenced: set[int] = set() + + def repl(match: re.Match) -> str: + index = int(match.group(1)) + if index < 1 or index > voice_count: + raise ValueError( + f"@Voice{index} does not match any connected voice ({voice_count} connected)." + ) + referenced.add(index) + return f"<|speaker:{index - 1}|>" + + rewritten = re.sub(r"(? list: + return [ + IO.Float.Input( + "temperature", + default=0.7, + min=0.0, + max=1.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Expressiveness. Higher values are more varied, lower values are more consistent.", + ), + IO.Float.Input( + "top_p", + default=0.7, + min=0.01, + max=1.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Diversity via nucleus sampling.", + ), + IO.Float.Input( + "speed", + default=1.0, + min=0.5, + max=2.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Speaking rate. 1.0 is normal, <1.0 slower, >1.0 faster.", + ), + IO.Float.Input( + "volume", + default=0.0, + min=-10.0, + max=10.0, + step=0.5, + display_mode=IO.NumberDisplay.slider, + tooltip="Volume adjustment in decibels. 0 is no change.", + ), + IO.Boolean.Input( + "normalize", + default=True, + tooltip="Normalize numbers and text for English and Chinese, " + "improving stability for numbers and dates.", + ), + ] + + +def _multi_speaker_inputs() -> list: + return [ + IO.Autogrow.Input( + "voices", + template=IO.Autogrow.TemplatePrefix( + IO.Custom(FISHAUDIO_VOICE).Input("voice"), + prefix="voice", + min=0, + max=5, + ), + tooltip="Voices for synthesis. Leave empty for the default voice. " + "With two or more voices, mark speaker changes in the text with @Voice1, @Voice2, etc.", + ), + *_tts_option_inputs(), + ] + + +class FishAudioVoiceSelector(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioVoiceSelector", + display_name="Fish Audio Voice Selector", + category="partner/audio/Fish Audio", + description="Select a voice from the Fish Audio library for text-to-speech generation.", + inputs=[ + IO.DynamicCombo.Input( + "voice", + options=[ + *(IO.DynamicCombo.Option(label, []) for _, label in FISHAUDIO_VOICES), + IO.DynamicCombo.Option( + "custom", + [ + IO.String.Input( + "voice_id", + default="", + tooltip="Voice model ID from fish.audio, e.g. the ID in " + "https://fish.audio/m//.", + ), + ], + ), + ], + tooltip="Choose a voice, or 'custom' to enter any fish.audio voice model ID.", + ), + ], + outputs=[ + IO.Custom(FISHAUDIO_VOICE).Output(display_name="voice"), + ], + is_api_node=False, + ) + + @classmethod + def execute(cls, voice: dict) -> IO.NodeOutput: + selected = voice["voice"] + if selected == "custom": + voice_id = voice["voice_id"].strip() + if not voice_id: + raise ValueError("Custom voice ID is empty.") + return IO.NodeOutput(voice_id) + voice_id = FISHAUDIO_VOICE_MAP.get(selected) + if not voice_id: + raise ValueError(f"Unknown voice: {selected}") + return IO.NodeOutput(voice_id) + + +class FishAudioTextToSpeech(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioTextToSpeech", + display_name="Fish Audio Text to Speech", + category="partner/audio/Fish Audio", + description="Convert text to speech. Supports emotion cues in the text " + "([happy], [whispering] on s2.1-pro; (happy) on s1) and multi-speaker dialogue " + "via @Voice1/@Voice2 tags with multiple connected voices.", + inputs=[ + IO.String.Input( + "text", + multiline=True, + default="", + tooltip="The text to convert to speech. With two or more voices connected, " + "mark speaker changes with @Voice1, @Voice2, etc.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option("s2.1-pro", _multi_speaker_inputs()), + IO.DynamicCombo.Option( + "s1", + [ + IO.Custom(FISHAUDIO_VOICE).Input( + "voice", + optional=True, + tooltip="Voice for synthesis. Leave unconnected for the default voice.", + ), + *_tts_option_inputs(), + ], + ), + ], + tooltip="Model to use for text-to-speech.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + ], + outputs=[ + IO.Audio.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["text"]), + expr=""" + ( + $t := widgets.text; + $type($t) = "string" + ? ( + $bytes := $length($t) + 2 * $count($match($t, /[^\\x00-\\x7F]/)); + {"type":"usd","usd": $bytes * 21.45 / 1000000, "format":{"approximate":true}} + ) + : {"type":"usd","usd": 0.02145, "format":{"approximate":true, "suffix":"/1K bytes"}} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + text: str, + model: dict, + seed: int, + ) -> IO.NodeOutput: + validate_string(text, field_name="text", min_length=1) + model_name = model["model"] + if model_name == "s1": + voices = [model["voice"]] if model.get("voice") else [] + else: + voices = [model["voices"][key] for key in model["voices"]] + rewritten, referenced = _rewrite_voice_tags(text, len(voices)) + if len(voices) >= 2: + missing = [i for i in range(1, len(voices) + 1) if i not in referenced] + if missing: + raise ValueError( + "With multiple voices, the text must mark speaker changes with tags for " + "each connected voice; missing: " + ", ".join(f"@Voice{i}" for i in missing) + ) + reference_id: str | list[str] | None = None + if len(voices) == 1: + reference_id = voices[0] + elif voices: + reference_id = voices + request = FishAudioTTSRequest( + text=rewritten, + reference_id=reference_id, + temperature=model["temperature"], + top_p=model["top_p"], + prosody=FishAudioProsody(speed=model["speed"], volume=model["volume"]), + normalize=model["normalize"], + ) + response = await sync_op_raw( + cls, + ApiEndpoint( + path="/proxy/fishaudio/v1/tts", + method="POST", + headers={"model": model_name}, + ), + data=request, + as_binary=True, + ) + return IO.NodeOutput(audio_bytes_to_audio_input(response)) + + +class FishAudioSpeechToText(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioSpeechToText", + display_name="Fish Audio Speech to Text", + category="partner/audio/Fish Audio", + description="Transcribe audio to text with automatic language detection.", + inputs=[ + IO.Audio.Input( + "audio", + tooltip="Audio to transcribe.", + ), + IO.String.Input( + "language", + default="", + tooltip="ISO 639-1 language hint (e.g. 'en', 'zh'). " + "The language is auto-detected regardless.", + ), + IO.Boolean.Input( + "precise_timestamps", + default=False, + tooltip="Return word-level timestamped segments.", + ), + ], + outputs=[ + IO.String.Output(id="text", display_name="text"), + IO.String.Output(id="language_code", display_name="language_code"), + IO.String.Output(id="segments_json", display_name="segments_json"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.00858,"format":{"approximate":true,"suffix":"/minute"}}""", + ), + ) + + @classmethod + async def execute( + cls, + audio: Input.Audio, + language: str, + precise_timestamps: bool, + ) -> IO.NodeOutput: + audio_data_np = audio_tensor_to_contiguous_ndarray(audio["waveform"]) + audio_bytes_io = audio_ndarray_to_bytesio(audio_data_np, audio["sample_rate"], "mp4", "aac") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/fishaudio/v1/asr", method="POST"), + response_model=FishAudioASRResponse, + data=FishAudioASRRequest( + language=language.strip() or None, + ignore_timestamps=not precise_timestamps, + ), + files={"audio": ("audio.mp4", audio_bytes_io, "audio/mp4")}, + content_type="multipart/form-data", + ) + segments_json = json.dumps( + [s.model_dump(exclude_none=True) for s in (response.segments or [])], + indent=2, + ) + return IO.NodeOutput(response.text or "", response.language_code or "", segments_json) + + +class FishAudioInstantVoiceClone(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioInstantVoiceClone", + display_name="Fish Audio Instant Voice Clone", + category="partner/audio/Fish Audio", + description="Create a private cloned voice from audio samples, instantly usable " + "for text-to-speech. Provide 1-20 recordings, 10-30 seconds each recommended, " + "under 270 seconds in total.", + inputs=[ + IO.Autogrow.Input( + "files", + template=IO.Autogrow.TemplatePrefix( + IO.Audio.Input("audio"), + prefix="audio", + min=1, + max=20, + ), + tooltip="Audio recordings for voice cloning.", + ), + IO.Boolean.Input( + "enhance_audio_quality", + default=True, + tooltip="Enhance reference audio quality before training.", + ), + ], + outputs=[ + IO.Custom(FISHAUDIO_VOICE).Output(display_name="voice"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge(expr="""{"type":"usd","usd":0}"""), + ) + + @classmethod + async def execute( + cls, + files: IO.Autogrow.Type, + enhance_audio_quality: bool, + ) -> IO.NodeOutput: + total_seconds = 0.0 + for key in files: + audio = files[key] + total_seconds += audio["waveform"].shape[-1] / audio["sample_rate"] + if total_seconds >= MAX_REFERENCE_AUDIO_SECONDS: + raise ValueError( + f"Total reference audio is {total_seconds:.0f} seconds; " + f"it must be under {MAX_REFERENCE_AUDIO_SECONDS} seconds." + ) + file_tuples: list[tuple[str, tuple[str, bytes, str]]] = [] + for key in files: + audio = files[key] + audio_data_np = audio_tensor_to_contiguous_ndarray(audio["waveform"]) + audio_bytes_io = audio_ndarray_to_bytesio(audio_data_np, audio["sample_rate"], "mp4", "aac") + file_tuples.append(("voices", (f"{key}.mp4", audio_bytes_io.getvalue(), "audio/mp4"))) + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/fishaudio/model", method="POST"), + response_model=FishAudioCreateModelResponse, + data=FishAudioCreateModelRequest( + title=str(uuid.uuid4()), + enhance_audio_quality=enhance_audio_quality, + ), + files=file_tuples, + content_type="multipart/form-data", + ) + return IO.NodeOutput(response.id) + + +class FishAudioExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + FishAudioVoiceSelector, + FishAudioTextToSpeech, + FishAudioSpeechToText, + FishAudioInstantVoiceClone, + ] + + +async def comfy_entrypoint() -> FishAudioExtension: + return FishAudioExtension() diff --git a/comfy_api_nodes/nodes_gemini.py b/comfy_api_nodes/nodes_gemini.py index 131590751..ebad67950 100644 --- a/comfy_api_nodes/nodes_gemini.py +++ b/comfy_api_nodes/nodes_gemini.py @@ -43,6 +43,7 @@ from comfy_api_nodes.util import ( download_url_to_image_tensor, download_url_to_video_output, get_number_of_images, + pad_images_to_common_channels, sync_op, tensor_to_base64_string, upload_audio_to_comfyapi, @@ -233,8 +234,8 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug "Try rephrasing your prompt or changing the response modality to 'IMAGE+TEXT' " "to see the model's reasoning." ) - return torch.zeros((1, 1024, 1024, 4)) - return torch.cat(image_tensors, dim=0) + return torch.zeros((1, 1024, 1024, 3)) + return torch.cat(pad_images_to_common_channels(image_tensors), dim=0) def get_text_from_interaction(interaction: GeminiInteraction) -> str: @@ -591,6 +592,7 @@ class GeminiNode(IO.ComfyNode): GEMINI_V2_MODELS: dict[str, str] = { + "Gemini 3.7 Flash": "gemini-3.7-flash", "Gemini 3.1 Pro": "gemini-3.1-pro-preview", "Gemini 3.5 Flash": "gemini-3.5-flash", "Gemini 3.1 Flash-Lite": "gemini-3.1-flash-lite-preview", @@ -693,6 +695,10 @@ class GeminiNodeV2(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "Gemini 3.7 Flash", + _gemini_text_model_inputs("MEDIUM", ["LOW", "MEDIUM", "HIGH"]), + ), IO.DynamicCombo.Option( "Gemini 3.5 Flash", _gemini_text_model_inputs("MEDIUM", ["MINIMAL", "LOW", "MEDIUM", "HIGH"]), @@ -738,6 +744,11 @@ class GeminiNodeV2(IO.ComfyNode): "usd": [0.00025, 0.0015], "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } } + : $contains($m, "3.7 flash") ? { + "type": "list_usd", + "usd": [0.00215, 0.01073], + "format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" } + } : $contains($m, "3.5 flash") ? { "type": "list_usd", "usd": [0.0015, 0.009], diff --git a/comfy_api_nodes/nodes_grok.py b/comfy_api_nodes/nodes_grok.py index 672a3e537..c58fb1a29 100644 --- a/comfy_api_nodes/nodes_grok.py +++ b/comfy_api_nodes/nodes_grok.py @@ -36,6 +36,26 @@ _GROK_VIDEO_MODEL_API_IDS = { "grok-imagine-video-1.5": "grok-imagine-video-1.5", } +_GROK_IMAGE_MODEL_API_IDS = { + "grok-imagine-image-2.0": "grok-imagine-image-2.0", +} + +_GROK_IMAGE_QUALITY_MODELS = {"grok-imagine-image-2.0"} + +_GROK_IMAGE_QUALITY_OPTIONS = ["medium", "low"] + +_GROK_IMAGE_EDIT_MAX_IMAGES = { + "grok-imagine-image-2.0": 3, + "grok-imagine-image-pro": 1, + "grok-imagine-image-quality": 3, + "grok-imagine-image": 3, +} + +_GROK_IMAGE_EDIT_ASPECT_RATIO_NEEDS_MULTIPLE = { + "grok-imagine-image-quality", + "grok-imagine-image", +} + _GROK_VOICE_OPTIONS = [ "none", "ara", @@ -106,19 +126,6 @@ def _normalize_grok_reference_prompt(prompt: str, total_images: int, voices: lis return prompt -def _extract_grok_price(response) -> float | None: - if response.usage and response.usage.cost_in_usd_ticks is not None: - return response.usage.cost_in_usd_ticks / 10_000_000_000 - return None - - -def _extract_grok_video_price(response) -> float | None: - price = _extract_grok_price(response) - if price is not None: - return price * 1.43 - return None - - class GrokImageNode(IO.ComfyNode): @classmethod @@ -132,6 +139,7 @@ class GrokImageNode(IO.ComfyNode): IO.Combo.Input( "model", options=[ + "grok-imagine-image-2.0", "grok-imagine-image-quality", "grok-imagine-image-pro", "grok-imagine-image", @@ -181,6 +189,12 @@ class GrokImageNode(IO.ComfyNode): "actual results are nondeterministic regardless of seed.", ), IO.Combo.Input("resolution", options=["1K", "2K"], optional=True), + IO.Combo.Input( + "quality", + options=_GROK_IMAGE_QUALITY_OPTIONS, + optional=True, + tooltip="Quality level, supported only by the grok-imagine-image-2.0 model.", + ), ], outputs=[ IO.Image.Output(), @@ -192,12 +206,15 @@ class GrokImageNode(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model", "number_of_images", "resolution"]), + depends_on=IO.PriceBadgeDepends(widgets=["model", "number_of_images", "resolution", "quality"]), expr=""" ( - $rate := widgets.model = "grok-imagine-image-quality" - ? (widgets.resolution = "1k" ? 0.05 : 0.07) - : ($contains(widgets.model, "pro") ? 0.07 : 0.02); + $is1k := widgets.resolution = "1k"; + $rate := widgets.model = "grok-imagine-image-2.0" + ? (widgets.quality = "low" ? ($is1k ? 0.04 : 0.06) : ($is1k ? 0.06 : 0.08)) + : (widgets.model = "grok-imagine-image-quality" + ? ($is1k ? 0.05 : 0.07) + : ($contains(widgets.model, "pro") ? 0.07 : 0.02)); {"type":"usd","usd": $rate * widgets.number_of_images} ) """, @@ -213,18 +230,20 @@ class GrokImageNode(IO.ComfyNode): number_of_images: int, seed: int, resolution: str = "1K", + quality: str = "medium", ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1) response = await sync_op( cls, ApiEndpoint(path="/proxy/xai/v1/images/generations", method="POST"), data=ImageGenerationRequest( - model=model, + model=_GROK_IMAGE_MODEL_API_IDS.get(model, model), prompt=prompt, aspect_ratio=aspect_ratio, n=number_of_images, seed=seed, resolution=resolution.lower(), + quality=quality if model in _GROK_IMAGE_QUALITY_MODELS else None, ), response_model=ImageGenerationResponse, ) @@ -255,7 +274,9 @@ _GROK_IMAGE_EDIT_ASPECT_RATIO_OPTIONS = [ ] -def _grok_image_edit_model_inputs(*, max_ref_images: int, with_aspect_ratio: bool): +def _grok_image_edit_model_inputs( + *, max_ref_images: int, with_aspect_ratio: bool, with_quality: bool = False, aspect_ratio_needs_multiple: bool = True +): inputs = [ IO.Autogrow.Input( "images", @@ -281,12 +302,18 @@ def _grok_image_edit_model_inputs(*, max_ref_images: int, with_aspect_ratio: boo display_mode=IO.NumberDisplay.number, ), ] + if with_quality: + inputs.append(IO.Combo.Input("quality", options=_GROK_IMAGE_QUALITY_OPTIONS)) if with_aspect_ratio: inputs.append( IO.Combo.Input( "aspect_ratio", options=_GROK_IMAGE_EDIT_ASPECT_RATIO_OPTIONS, - tooltip="Only allowed when multiple images are connected.", + tooltip=( + "Only allowed when multiple images are connected." + if aspect_ratio_needs_multiple + else "Aspect ratio of the edited image." + ), ) ) return inputs @@ -451,6 +478,15 @@ class GrokImageEditNodeV2(IO.ComfyNode): IO.DynamicCombo.Input( "model", options=[ + IO.DynamicCombo.Option( + "grok-imagine-image-2.0", + _grok_image_edit_model_inputs( + max_ref_images=3, + with_aspect_ratio=True, + with_quality=True, + aspect_ratio_needs_multiple=False, + ), + ), IO.DynamicCombo.Option( "grok-imagine-image-quality", _grok_image_edit_model_inputs(max_ref_images=3, with_aspect_ratio=True), @@ -488,18 +524,23 @@ class GrokImageEditNodeV2(IO.ComfyNode): is_api_node=True, price_badge=IO.PriceBadge( depends_on=IO.PriceBadgeDepends( - widgets=["model", "model.resolution", "model.number_of_images"], + widgets=["model", "model.resolution", "model.number_of_images", "model.quality"], ), expr=""" ( - $isQualityModel := widgets.model = "grok-imagine-image-quality"; + $is20 := widgets.model = "grok-imagine-image-2.0"; $isPro := $contains(widgets.model, "pro"); $res := $lookup(widgets, "model.resolution"); $n := $lookup(widgets, "model.number_of_images"); - $rate := $isQualityModel - ? ($res = "1k" ? 0.05 : 0.07) - : ($isPro ? 0.07 : 0.02); - $base := $isQualityModel ? 0.01 : 0.002; + $is1k := $res = "1k"; + $rate := $is20 + ? ($lookup(widgets, "model.quality") = "low" + ? ($is1k ? 0.04 : 0.06) + : ($is1k ? 0.06 : 0.08)) + : (widgets.model = "grok-imagine-image-quality" + ? ($is1k ? 0.05 : 0.07) + : ($isPro ? 0.07 : 0.02)); + $base := ($is20 or widgets.model = "grok-imagine-image-quality") ? 0.01 : 0.002; $output := $rate * $n; $isPro ? {"type":"usd","usd": $base + $output} @@ -525,13 +566,15 @@ class GrokImageEditNodeV2(IO.ComfyNode): image_tensors: list[Input.Image] = [t for t in images_dict.values() if t is not None] n_images = sum(get_number_of_images(t) for t in image_tensors) + max_images = _GROK_IMAGE_EDIT_MAX_IMAGES.get(model_id, 3) if n_images < 1: raise ValueError("At least one image is required for editing.") - if model_id == "grok-imagine-image-pro" and n_images > 1: - raise ValueError("The pro model supports only 1 input image.") - if model_id != "grok-imagine-image-pro" and n_images > 3: - raise ValueError("A maximum of 3 input images is supported.") - if aspect_ratio != "auto" and n_images == 1: + if n_images > max_images: + raise ValueError( + f"The {model_id} model supports at most {max_images} input " + f"image{'s' if max_images > 1 else ''}; {n_images} are connected." + ) + if aspect_ratio != "auto" and model_id in _GROK_IMAGE_EDIT_ASPECT_RATIO_NEEDS_MULTIPLE and n_images == 1: raise ValueError( "Custom aspect ratio is only allowed when multiple images are connected to the image input." ) @@ -547,7 +590,7 @@ class GrokImageEditNodeV2(IO.ComfyNode): cls, ApiEndpoint(path="/proxy/xai/v1/images/edits", method="POST"), data=ImageEditRequest( - model=model_id, + model=_GROK_IMAGE_MODEL_API_IDS.get(model_id, model_id), images=[ InputUrlObject(url=f"data:image/png;base64,{tensor_to_base64_string(i)}") for i in flat_tensors ], @@ -556,6 +599,7 @@ class GrokImageEditNodeV2(IO.ComfyNode): n=number_of_images, seed=seed, aspect_ratio=None if aspect_ratio == "auto" else aspect_ratio, + quality=model.get("quality") if model_id in _GROK_IMAGE_QUALITY_MODELS else None, ), response_model=ImageGenerationResponse, ) @@ -690,7 +734,6 @@ class GrokVideoNode(IO.ComfyNode): ApiEndpoint(path=f"/proxy/xai/v1/videos/{initial_response.request_id}"), status_extractor=lambda r: r.status if r.status is not None else "complete", response_model=VideoStatusResponse, - price_extractor=_extract_grok_video_price if model == "grok-imagine-video-1.5" else _extract_grok_price, ) return IO.NodeOutput(await download_url_to_video_output(response.video.url)) @@ -768,7 +811,6 @@ class GrokVideoEditNode(IO.ComfyNode): ApiEndpoint(path=f"/proxy/xai/v1/videos/{initial_response.request_id}"), status_extractor=lambda r: r.status if r.status is not None else "complete", response_model=VideoStatusResponse, - price_extractor=_extract_grok_price, ) return IO.NodeOutput(await download_url_to_video_output(response.video.url)) @@ -965,7 +1007,6 @@ class GrokVideoReferenceNode(IO.ComfyNode): ApiEndpoint(path=f"/proxy/xai/v1/videos/{initial_response.request_id}"), status_extractor=lambda r: r.status if r.status is not None else "complete", response_model=VideoStatusResponse, - price_extractor=_extract_grok_video_price, ) return IO.NodeOutput(await download_url_to_video_output(response.video.url)) @@ -1070,7 +1111,6 @@ class GrokVideoExtendNode(IO.ComfyNode): ApiEndpoint(path=f"/proxy/xai/v1/videos/{initial_response.request_id}"), status_extractor=lambda r: r.status if r.status is not None else "complete", response_model=VideoStatusResponse, - price_extractor=_extract_grok_video_price, ) return IO.NodeOutput(await download_url_to_video_output(response.video.url)) diff --git a/comfy_api_nodes/nodes_hitpaw.py b/comfy_api_nodes/nodes_hitpaw.py index 062d3cf1d..40da2bfc1 100644 --- a/comfy_api_nodes/nodes_hitpaw.py +++ b/comfy_api_nodes/nodes_hitpaw.py @@ -169,14 +169,12 @@ class HitPawGeneralImageEnhance(IO.ComfyNode): ) if initial_res.code != 200: raise ValueError(f"Task creation failed with code {initial_res.code}: {initial_res.message}") - request_price = initial_res.data.consume_coins / 1000 final_response = await poll_op( cls, ApiEndpoint(path="/proxy/hitpaw/api/task-status", method="POST"), data=TaskCreateDataResponse(job_id=initial_res.data.job_id), response_model=TaskStatusResponse, status_extractor=lambda x: x.data.status, - price_extractor=lambda x: request_price, poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_image_tensor(final_response.data.res_url)) @@ -312,7 +310,6 @@ class HitPawVideoEnhance(IO.ComfyNode): wait_label="Creating task", final_label_on_success="Task created", ) - request_price = initial_res.data.consume_coins / 1000 if initial_res.code != 200: raise ValueError(f"Task creation failed with code {initial_res.code}: {initial_res.message}") final_response = await poll_op( @@ -321,7 +318,6 @@ class HitPawVideoEnhance(IO.ComfyNode): data=TaskStatusPollRequest(job_id=initial_res.data.job_id), response_model=TaskStatusResponse, status_extractor=lambda x: x.data.status, - price_extractor=lambda x: request_price, poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_video_output(final_response.data.res_url)) diff --git a/comfy_api_nodes/nodes_ideogram.py b/comfy_api_nodes/nodes_ideogram.py index 252617b2c..2acf77b88 100644 --- a/comfy_api_nodes/nodes_ideogram.py +++ b/comfy_api_nodes/nodes_ideogram.py @@ -531,7 +531,7 @@ class IdeogramPImage(IO.ComfyNode): def define_schema(cls): return IO.Schema( node_id="IdeogramPImage", - display_name="Ideogram P-Image", + display_name="Ideogram & Pruna P-Image", category="partner/image/Ideogram", description="Generates images using P-Image, Ideogram's fast text-to-image model. " "Strong typography and photorealism; " diff --git a/comfy_api_nodes/nodes_kling.py b/comfy_api_nodes/nodes_kling.py index bbeabead1..1bbec1d1f 100644 --- a/comfy_api_nodes/nodes_kling.py +++ b/comfy_api_nodes/nodes_kling.py @@ -1866,7 +1866,7 @@ class KlingImageGenerationNode(IO.ComfyNode): tooltip="Subject reference similarity", advanced=True, ), - IO.Combo.Input("model_name", options=["kling-v3", "kling-v2"]), + IO.Combo.Input("model_name", options=["kling-v3"]), IO.Combo.Input( "aspect_ratio", options=[i.value for i in KlingImageGenAspectRatio], @@ -1902,13 +1902,8 @@ class KlingImageGenerationNode(IO.ComfyNode): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model_name", "n"]), - expr=""" - ( - $base := $contains(widgets.model_name,"kling-v3") ? 0.028 : 0.014; - {"type":"usd","usd": $base * widgets.n} - ) - """, + depends_on=IO.PriceBadgeDepends(widgets=["n"]), + expr="""{"type":"usd","usd": 0.028 * widgets.n}""", ), ) diff --git a/comfy_api_nodes/nodes_ltxv.py b/comfy_api_nodes/nodes_ltxv.py index 878e04b4e..44723dde2 100644 --- a/comfy_api_nodes/nodes_ltxv.py +++ b/comfy_api_nodes/nodes_ltxv.py @@ -6,8 +6,12 @@ from typing_extensions import override from comfy_api.latest import IO, ComfyExtension, Input, InputImpl from comfy_api_nodes.util import ( ApiEndpoint, + download_url_to_video_output, get_number_of_images, + poll_op, + sync_op, sync_op_raw, + upload_audio_to_comfyapi, upload_images_to_comfyapi, validate_string, ) @@ -17,6 +21,11 @@ MODELS_MAP = { "LTX-2 (Fast)": "ltx-2-fast", } +V25_MODELS_MAP = { + "LTX-2.5 (Fast)": "ltx-2-5-fast", + "LTX-2.5 (Pro)": "ltx-2-5-pro", +} + class ExecuteTaskRequest(BaseModel): prompt: str = Field(...) @@ -26,6 +35,48 @@ class ExecuteTaskRequest(BaseModel): fps: int | None = Field(25) generate_audio: bool | None = Field(True) image_uri: str | None = Field(None) + last_frame_uri: str | None = Field(None) + + +class AudioToVideoRequest(BaseModel): + prompt: str = Field(...) + model: str = Field(...) + resolution: str = Field(...) + audio_uri: str = Field(...) + image_uri: str | None = Field(None) + + +class Ltx25SubmitResponse(BaseModel): + id: str = Field(...) + + +class Ltx25JobResult(BaseModel): + video_url: str | None = Field(None) + + +class Ltx25JobStatusResponse(BaseModel): + id: str = Field(...) + status: str = Field(...) + result: Ltx25JobResult | None = Field(None) + + +async def _v25_submit_and_poll(cls: type[IO.ComfyNode], route: str, data: BaseModel) -> IO.NodeOutput: + submit = await sync_op( + cls, + ApiEndpoint(f"/proxy/ltx/v2/{route}", "POST"), + response_model=Ltx25SubmitResponse, + data=data, + max_retries=1, + ) + job = await poll_op( + cls, + ApiEndpoint(f"/proxy/ltx/v2/{route}/{submit.id}"), + response_model=Ltx25JobStatusResponse, + status_extractor=lambda r: r.status, + ) + if not job.result or not job.result.video_url: + raise RuntimeError(f"LTX job {job.id} completed without a video URL.") + return IO.NodeOutput(await download_url_to_video_output(job.result.video_url, cls=cls)) PRICE_BADGE = IO.PriceBadge( @@ -43,6 +94,128 @@ PRICE_BADGE = IO.PriceBadge( """, ) +V25_PRICE_BADGE = IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.duration", "model.resolution"]), + expr=""" + ( + $prices := { + "ltx-2.5 (fast)": { + "1280x720":0.1287,"720x1280":0.1287, + "1920x1080":0.1859,"1080x1920":0.1859, + "2560x1440":0.2717,"1440x2560":0.2717, + "3840x2160":0.429,"2160x3840":0.429 + }, + "ltx-2.5 (pro)": { + "1280x720":0.1716,"720x1280":0.1716, + "1920x1080":0.2431,"1080x1920":0.2431 + } + }; + $model := $lookup(widgets, "model"); + $table := $type($model) = "string" ? $lookup($prices, $model) : undefined; + $res := $lookup(widgets, "model.resolution"); + $pps := $type($table) = "object" and $type($res) = "string" ? $lookup($table, $res) : undefined; + $durRaw := $lookup(widgets, "model.duration"); + $dur := $type($durRaw) in ["string", "number"] ? $number($durRaw) : undefined; + $type($pps) = "number" and $type($dur) = "number" + ? {"type":"usd","usd": $pps * $dur} + : undefined + ) + """, +) + +V25_A2V_PRICE_BADGE = IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model"]), + expr=""" + ( + $rates := {"ltx-2.5 (fast)":0.1859, "ltx-2.5 (pro)":0.2431}; + $model := $lookup(widgets, "model"); + $rate := $type($model) = "string" ? $lookup($rates, $model) : undefined; + $type($rate) = "number" + ? {"type":"usd","usd": $rate, "format":{"suffix":"/second"}} + : undefined + ) + """, +) + + +def _v25_generation_inputs( + durations: list[str], resolutions: list[str], fps_options: list[str], tooltip: str | None +) -> list: + return [ + IO.Combo.Input( + "duration", + options=durations, + default="8", + tooltip=tooltip, + ), + IO.Combo.Input( + "resolution", + options=resolutions, + default="1920x1080", + ), + IO.Combo.Input("fps", options=fps_options, default="25"), + IO.Boolean.Input( + "generate_audio", + default=True, + tooltip="When true, the generated video will include AI-generated audio matching the scene.", + advanced=True, + ), + ] + + +def _v25_model_combo() -> IO.DynamicCombo.Input: + return IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "LTX-2.5 (Fast)", + _v25_generation_inputs( + ["2", "3", "4", "5", "6", "8", "10", "12", "14", "16", "18", "20"], + [ + "1280x720", + "720x1280", + "1920x1080", + "1080x1920", + "2560x1440", + "1440x2560", + "3840x2160", + "2160x3840", + ], + ["24", "25", "48", "50"], + "Video duration in seconds. Durations over 10s require a 720p/1080p resolution and 24/25 FPS.", + ), + ), + IO.DynamicCombo.Option( + "LTX-2.5 (Pro)", + _v25_generation_inputs( + ["2", "3", "4", "5", "6", "8", "10"], + ["1280x720", "720x1280", "1920x1080", "1080x1920"], + ["24", "25", "50"], + "Video duration in seconds.", + ), + ), + ], + ) + + +def _v25_seed_input() -> IO.Int.Input: + return IO.Int.Input( + "seed", + default=42, + min=0, + max=0xFFFFFFFF, + control_after_generate=True, + tooltip="Seed to determine if node should re-run; " + "actual results are nondeterministic regardless of seed.", + ) + + +def _v25_validate_settings(model: dict) -> None: + if int(model["duration"]) > 10 and ( + int(model["fps"]) > 25 or model["resolution"] in ("2560x1440", "1440x2560", "3840x2160", "2160x3840") + ): + raise ValueError("Durations over 10s require a 720p or 1080p resolution and 24/25 FPS.") + class TextToVideoNode(IO.ComfyNode): @classmethod @@ -86,6 +259,7 @@ class TextToVideoNode(IO.ComfyNode): IO.Hidden.unique_id, ], is_api_node=True, + is_deprecated=True, price_badge=PRICE_BADGE, ) @@ -164,6 +338,7 @@ class ImageToVideoNode(IO.ComfyNode): IO.Hidden.unique_id, ], is_api_node=True, + is_deprecated=True, price_badge=PRICE_BADGE, ) @@ -203,12 +378,217 @@ class ImageToVideoNode(IO.ComfyNode): return IO.NodeOutput(InputImpl.VideoFromFile(BytesIO(response))) +class Ltx25TextToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25TextToVideo", + display_name="LTX 2.5 Text To Video", + category="partner/video/LTXV", + description="Professional-quality videos with customizable duration and resolution.", + inputs=[ + _v25_model_combo(), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + model: dict, + prompt: str, + seed: int = 42, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + _v25_validate_settings(model) + return await _v25_submit_and_poll( + cls, + "text-to-video", + ExecuteTaskRequest( + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + duration=int(model["duration"]), + resolution=model["resolution"], + fps=int(model["fps"]), + generate_audio=model["generate_audio"], + ), + ) + + +class Ltx25ImageToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25ImageToVideo", + display_name="LTX 2.5 Image To Video", + category="partner/video/LTXV", + description="Professional-quality videos with customizable duration and resolution based on start image.", + inputs=[ + IO.Image.Input("image", tooltip="First frame to be used for the video."), + _v25_model_combo(), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + IO.Image.Input( + "last_frame", + optional=True, + tooltip="Last frame to be used for the video.", + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + image: Input.Image, + model: dict, + prompt: str, + seed: int = 42, + last_frame: Input.Image | None = None, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + _v25_validate_settings(model) + if get_number_of_images(image) != 1: + raise ValueError("Currently only one input image is supported.") + last_frame_uri = None + if last_frame is not None: + if get_number_of_images(last_frame) != 1: + raise ValueError("Currently only one last frame image is supported.") + last_frame_uri = (await upload_images_to_comfyapi(cls, last_frame, max_images=1, mime_type="image/png"))[0] + return await _v25_submit_and_poll( + cls, + "image-to-video", + ExecuteTaskRequest( + image_uri=(await upload_images_to_comfyapi(cls, image, max_images=1, mime_type="image/png"))[0], + last_frame_uri=last_frame_uri, + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + duration=int(model["duration"]), + resolution=model["resolution"], + fps=int(model["fps"]), + generate_audio=model["generate_audio"], + ), + ) + + +class Ltx25AudioToVideoNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="LtxApi25AudioToVideo", + display_name="LTX 2.5 Audio To Video", + category="partner/video/LTXV", + description="Generate a video driven by an audio track, with an optional first frame image.", + inputs=[ + IO.Audio.Input( + "audio", + tooltip="Audio track driving the video. Its length (2-20 seconds) sets the video duration.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "LTX-2.5 (Fast)", + [IO.Combo.Input("resolution", options=["1920x1080", "1080x1920"])], + ), + IO.DynamicCombo.Option( + "LTX-2.5 (Pro)", + [IO.Combo.Input("resolution", options=["1920x1080", "1080x1920"])], + ), + ], + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + ), + _v25_seed_input(), + IO.Image.Input( + "image", + optional=True, + tooltip="Optional first frame to be used for the video.", + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=V25_A2V_PRICE_BADGE, + ) + + @classmethod + async def execute( + cls, + audio: Input.Audio, + model: dict, + prompt: str, + seed: int = 42, + image: Input.Image | None = None, + ) -> IO.NodeOutput: + validate_string(prompt, min_length=1, max_length=10000) + audio_duration = audio["waveform"].shape[-1] / audio["sample_rate"] + if not 2 <= audio_duration <= 20: + raise ValueError(f"Audio duration must be between 2 and 20 seconds, got {audio_duration:.1f}s.") + image_uri = None + if image is not None: + if get_number_of_images(image) != 1: + raise ValueError("Currently only one input image is supported.") + image_uri = (await upload_images_to_comfyapi(cls, image, max_images=1, mime_type="image/png"))[0] + return await _v25_submit_and_poll( + cls, + "audio-to-video", + AudioToVideoRequest( + prompt=prompt, + model=V25_MODELS_MAP[model["model"]], + resolution=model["resolution"], + audio_uri=await upload_audio_to_comfyapi(cls, audio), + image_uri=image_uri, + ), + ) + + class LtxvApiExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: return [ TextToVideoNode, ImageToVideoNode, + Ltx25TextToVideoNode, + Ltx25ImageToVideoNode, + Ltx25AudioToVideoNode, ] diff --git a/comfy_api_nodes/nodes_magnific.py b/comfy_api_nodes/nodes_magnific.py index 4ce4735df..6e86fa795 100644 --- a/comfy_api_nodes/nodes_magnific.py +++ b/comfy_api_nodes/nodes_magnific.py @@ -30,31 +30,6 @@ from comfy_api_nodes.util import ( validate_image_dimensions, ) -_EUR_TO_USD = 1.19 - - -def _tier_price_eur(megapixels: float) -> float: - """Price in EUR for a single Magnific upscaling step based on input megapixels.""" - if megapixels <= 1.3: - return 0.143 - if megapixels <= 3.0: - return 0.286 - if megapixels <= 6.4: - return 0.429 - return 1.716 - - -def _calculate_magnific_upscale_price_usd(width: int, height: int, scale: int) -> float: - """Calculate total Magnific upscale price in USD for given input dimensions and scale factor.""" - num_steps = int(math.log2(scale)) - total_eur = 0.0 - pixels = width * height - for _ in range(num_steps): - total_eur += _tier_price_eur(pixels / 1_000_000) - pixels *= 4 - return round(total_eur * _EUR_TO_USD, 2) - - class MagnificImageUpscalerCreativeNode(IO.ComfyNode): @classmethod def define_schema(cls): @@ -203,10 +178,6 @@ class MagnificImageUpscalerCreativeNode(IO.ComfyNode): f"Use a smaller input image or lower scale factor." ) - final_height, final_width = get_image_dimensions(image) - actual_scale = int(scale_factor.rstrip("x")) - price_usd = _calculate_magnific_upscale_price_usd(final_width, final_height, actual_scale) - initial_res = await sync_op( cls, ApiEndpoint(path="/proxy/freepik/v1/ai/image-upscaler", method="POST"), @@ -228,7 +199,6 @@ class MagnificImageUpscalerCreativeNode(IO.ComfyNode): ApiEndpoint(path=f"/proxy/freepik/v1/ai/image-upscaler/{initial_res.task_id}"), response_model=TaskResponse, status_extractor=lambda x: x.status, - price_extractor=lambda _: price_usd, poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_image_tensor(final_response.generated[0])) @@ -367,9 +337,6 @@ class MagnificImageUpscalerPreciseV2Node(IO.ComfyNode): f"Use a smaller input image or lower scale factor." ) - final_height, final_width = get_image_dimensions(image) - price_usd = _calculate_magnific_upscale_price_usd(final_width, final_height, requested_scale) - initial_res = await sync_op( cls, ApiEndpoint(path="/proxy/freepik/v1/ai/image-upscaler-precision-v2", method="POST"), @@ -388,7 +355,6 @@ class MagnificImageUpscalerPreciseV2Node(IO.ComfyNode): ApiEndpoint(path=f"/proxy/freepik/v1/ai/image-upscaler-precision-v2/{initial_res.task_id}"), response_model=TaskResponse, status_extractor=lambda x: x.status, - price_extractor=lambda _: price_usd, poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_image_tensor(final_response.generated[0])) diff --git a/comfy_api_nodes/nodes_minimax.py b/comfy_api_nodes/nodes_minimax.py index 3c1d29257..de3895221 100644 --- a/comfy_api_nodes/nodes_minimax.py +++ b/comfy_api_nodes/nodes_minimax.py @@ -3,12 +3,14 @@ from typing import Optional import torch from typing_extensions import override -from comfy_api.latest import IO, ComfyExtension +from comfy_api.latest import IO, ComfyExtension, Input from comfy_api_nodes.apis.minimax import ( Hailuo03AudioContent, Hailuo03AudioContentUrl, + Hailuo03ContextIRRequest, Hailuo03ImageContent, Hailuo03ImageContentUrl, + Hailuo03RegenerationRequest, Hailuo03TaskCreationRequest, Hailuo03TaskCreationResponse, Hailuo03TaskQueryResponse, @@ -456,6 +458,9 @@ HAILUO_03_QUERY_ENDPOINT = "/proxy/minimax/v2/query/video_generation" # + /{tas HAILUO_03_MODELS = {"MiniMax H3": "MiniMax-H3"} HAILUO_03_FAILED_STATUSES = ["failed", "cancelled", "expired"] +HAILUO_03_CONTEXT_IR_ENDPOINT = "/proxy/minimax/v2/h3_context_ir" +HAILUO_03_REGENERATION_ENDPOINT = "/proxy/minimax/v2/video_regeneration" + def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = True): inputs = [ @@ -487,10 +492,10 @@ def _hailuo03_model_inputs(include_ratio: bool = True, allow_adaptive: bool = Tr IO.Int.Input( "duration", default=5, - min=5, + min=4, max=15, step=1, - tooltip="Duration of the output video in seconds (5-15).", + tooltip="Duration of the output video in seconds (4-15).", display_mode=IO.NumberDisplay.slider, ) ) @@ -939,6 +944,592 @@ class MinimaxHailuo03ReferenceNode(IO.ComfyNode): ) +class MinimaxHailuo03ContextIRNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03ContextIRNode", + display_name="MiniMax H3 Context IR (Prompt Enhancer)", + category="partner/video/MiniMax", + description="Analyze text and media context with MiniMax H3 Context IR and produce an enhanced, " + "structured video prompt. Feed the output into the prompt of a MiniMax H3 video node and attach " + "the same media there in the same order, because the enhanced prompt refers to the attached " + "media by position.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "MiniMax H3", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Description of the video you intend to generate.", + ), + IO.Int.Input( + "duration", + default=5, + min=4, + max=15, + step=1, + tooltip="Duration of the video you intend to generate, in seconds (4-15).", + display_mode=IO.NumberDisplay.slider, + ), + IO.Combo.Input( + "ratio", + options=["adaptive", "16:9", "4:3", "1:1", "3:4", "9:16", "21:9"], + default="adaptive", + tooltip="Aspect ratio of the video you intend to generate. 'adaptive' " + "requires at least one image, video, or audio input.", + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image_1", + "image_2", + "image_3", + "image_4", + "image_5", + "image_6", + "image_7", + "image_8", + "image_9", + ], + min=0, + ), + tooltip="Subject or style reference images, referred to in the prompt " + "as 'Image 1'..'Image 9' in connection order. Up to 9 images.", + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=["video_1", "video_2", "video_3"], + min=0, + ), + tooltip="Motion or scene reference videos, referred to in the prompt " + "as 'Video 1'..'Video 3' in connection order. Up to 3 videos, " + "2-15 seconds each, 15 seconds in total.", + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=["audio_1", "audio_2", "audio_3"], + min=0, + ), + tooltip="Audio references, referred to in the prompt as " + "'Audio 1'..'Audio 3' in connection order. Up to 3 clips, " + "2-15 seconds each, 15 seconds in total. Cannot be used without " + "a reference image or video.", + ), + ], + ) + ], + tooltip="Model to use for prompt enhancement.", + ), + IO.Image.Input( + "first_frame", + tooltip="First frame of the video you intend to generate. Cannot be combined with " + "reference media.", + optional=True, + ), + IO.Image.Input( + "last_frame", + tooltip="Last frame of the video you intend to generate. Cannot be combined with " + "reference media.", + optional=True, + ), + ], + outputs=[ + IO.String.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + inputs=["first_frame", "last_frame"], + input_groups=["model.reference_images", "model.reference_videos", "model.reference_audios"], + ), + expr=""" + ( + $imgsRaw := $lookup(inputGroups, "model.reference_images"); + $imgs := $imgsRaw ? $imgsRaw : 0; + $vidsRaw := $lookup(inputGroups, "model.reference_videos"); + $vids := $vidsRaw ? $vidsRaw : 0; + $audsRaw := $lookup(inputGroups, "model.reference_audios"); + $auds := $audsRaw ? $audsRaw : 0; + $frames := (inputs.first_frame.connected ? 1 : 0) + (inputs.last_frame.connected ? 1 : 0); + ($imgs + $vids + $auds) > 0 + ? {"type": "range_usd", "min_usd": 0.06, "max_usd": 0.11, "format": {"approximate": true}} + : $frames > 0 + ? {"type": "usd", "usd": 0.05, "format": {"approximate": true}} + : {"type": "usd", "usd": 0.02, "format": {"approximate": true}} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + first_frame: torch.Tensor | None = None, + last_frame: torch.Tensor | None = None, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + + reference_images = {k: v for k, v in (model.get("reference_images") or {}).items() if v is not None} + reference_videos = {k: v for k, v in (model.get("reference_videos") or {}).items() if v is not None} + reference_audios = {k: v for k, v in (model.get("reference_audios") or {}).items() if v is not None} + has_frames = first_frame is not None or last_frame is not None + has_references = bool(reference_images) or bool(reference_videos) or bool(reference_audios) + if has_frames and has_references: + raise ValueError( + "First/last frame and reference media are mutually exclusive. Use frames for an " + "image-to-video prompt, or reference media for a reference-to-video prompt." + ) + if reference_audios and not reference_images and not reference_videos: + raise ValueError("Reference audio cannot be used without a reference image or video.") + if not has_frames and not has_references and model["ratio"] == "adaptive": + raise ValueError( + "Ratio 'adaptive' is not supported for text-only requests; select an explicit aspect ratio." + ) + + for frame in (first_frame, last_frame): + if frame is not None: + validate_image_aspect_ratio(frame, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(frame, min_width=256, min_height=256) + for image in reference_images.values(): + validate_image_aspect_ratio(image, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(image, min_width=256, min_height=256) + + total_video_duration = 0.0 + for i, video in enumerate(reference_videos.values(), 1): + try: + fps = float(video.get_frame_rate()) + except Exception: + fps = 0.0 + if fps and not (23.9 <= fps <= 60.5): + raise ValueError(f"Reference video {i} is {fps:.2f} FPS. Supported range is 23.976-60 FPS.") + try: + dur = video.get_duration() + except Exception: + continue + if dur < 1.8: + raise ValueError(f"Reference video {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_video_duration += dur + if total_video_duration > 15.1: + raise ValueError( + f"Total reference video duration is {total_video_duration:.1f}s. Maximum is 15 seconds." + ) + + total_audio_duration = 0.0 + for i, audio in enumerate(reference_audios.values(), 1): + dur = int(audio["waveform"].shape[-1]) / int(audio["sample_rate"]) + if dur < 1.8: + raise ValueError(f"Reference audio {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_audio_duration += dur + if total_audio_duration > 15.1: + raise ValueError( + f"Total reference audio duration is {total_audio_duration:.1f}s. Maximum is 15 seconds." + ) + + content: list = [Hailuo03TextContent(text=model["prompt"])] + if first_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, first_frame, max_images=1, wait_label="Uploading first frame" + ) + )[0], + ), + role="first_frame", + ) + ) + if last_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, last_frame, max_images=1, wait_label="Uploading last frame" + ) + )[0], + ), + role="last_frame", + ) + ) + for i, image in enumerate(reference_images.values(), 1): + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, image, max_images=1, wait_label=f"Uploading image {i}" + ) + )[0], + ), + role="reference_image", + ) + ) + for i, video in enumerate(reference_videos.values(), 1): + content.append( + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, video, wait_label=f"Uploading video {i}"), + ), + ) + ) + for audio in reference_audios.values(): + content.append( + Hailuo03AudioContent( + audio_url=Hailuo03AudioContentUrl( + url=await upload_audio_to_comfyapi( + cls, + audio, + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ), + ) + ) + + response = await sync_op( + cls, + ApiEndpoint(path=HAILUO_03_CONTEXT_IR_ENDPOINT, method="POST"), + response_model=Hailuo03TaskCreationResponse, + data=Hailuo03ContextIRRequest( + model=HAILUO_03_MODELS[model["model"]], + content=content, + duration=model["duration"], + ratio=None if model["ratio"] == "adaptive" else model["ratio"], + ), + ) + task_result = await poll_op( + cls, + ApiEndpoint(path=f"{HAILUO_03_QUERY_ENDPOINT}/{response.task_id}"), + response_model=Hailuo03TaskQueryResponse, + status_extractor=lambda r: r.task.status, + failed_statuses=HAILUO_03_FAILED_STATUSES, + poll_interval=5, + ) + prompt = task_result.task.content.prompt if task_result.task.content else None + if not prompt: + raise Exception(f"No enhanced prompt in the response: {task_result.model_dump()}") + return IO.NodeOutput(prompt) + + +class MinimaxHailuo03RegenerateNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="MinimaxHailuo03RegenerateNode", + display_name="MiniMax H3 Regenerate to 2K", + category="partner/video/MiniMax", + description="Re-render a MiniMax H3 768P output at 2K resolution. Connect the unmodified 768P " + "video and the exact prompt used to generate it; if the original generation used first/last " + "frames or reference media, attach the same inputs.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "MiniMax H3", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="The exact prompt used to generate the source video.", + ), + IO.Combo.Input( + "resolution", + options=["2K"], + tooltip="Resolution to re-render the source video at.", + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[ + "image_1", + "image_2", + "image_3", + "image_4", + "image_5", + "image_6", + "image_7", + "image_8", + "image_9", + ], + min=0, + ), + tooltip="Reference images from the original generation, in the same " + "order. Up to 9 images.", + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=["video_1", "video_2", "video_3"], + min=0, + ), + tooltip="Reference videos from the original generation, in the same " + "order. Up to 3 videos, 2-15 seconds each, 15 seconds in total.", + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=["audio_1", "audio_2", "audio_3"], + min=0, + ), + tooltip="Audio references from the original generation, in the same " + "order. Up to 3 clips, 2-15 seconds each, 15 seconds in total. " + "Cannot be used without a reference image or video.", + ), + ], + ) + ], + tooltip="Model to use for video regeneration.", + ), + IO.Video.Input( + "video", + tooltip="The MiniMax H3 768P output video to re-render. Connect the unmodified output " + "of a MiniMax H3 video node (24 FPS, 4-15 seconds). 2K outputs cannot be used.", + ), + IO.Image.Input( + "first_frame", + tooltip="First frame image from the original generation, if one was used.", + optional=True, + ), + IO.Image.Input( + "last_frame", + tooltip="Last frame image from the original generation, if one was used.", + optional=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AIGC watermark to the video.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type": "usd", "usd": 0.0715, "format": {"suffix": "/second"}}""", + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + video: Input.Video, + watermark: bool, + first_frame: torch.Tensor | None = None, + last_frame: torch.Tensor | None = None, + ) -> IO.NodeOutput: + validate_string(model["prompt"], strip_whitespace=True, min_length=1) + + try: + fps = float(video.get_frame_rate()) + except Exception: + fps = 0.0 + if fps and not (23.9 <= fps <= 24.1): + raise ValueError( + f"The source video is {fps:.2f} FPS. Regeneration accepts unmodified MiniMax H3 768P " + "outputs, which are 24 FPS." + ) + try: + width, height = video.get_dimensions() + except Exception: + width = height = 0 + if width and height and (width % 32 or height % 32 or width * height > 1_032_192): + raise ValueError( + f"The source video is {width}x{height}. Regeneration accepts MiniMax H3 768P outputs " + "(width and height divisible by 32, at most 1,032,192 total pixels); 2K outputs cannot " + "be used as a source." + ) + try: + frame_count = video.get_frame_count() + except Exception: + frame_count = 0 + if frame_count and (frame_count < 107 or frame_count > 362 or (frame_count - 107) % 17): + raise ValueError( + f"The source video has {frame_count} frames. Regeneration accepts unmodified " + "MiniMax H3 outputs, whose length is 107 to 362 frames in steps of 17 " + "(4 to 15 seconds at 24 FPS)." + ) + + reference_images = {k: v for k, v in (model.get("reference_images") or {}).items() if v is not None} + reference_videos = {k: v for k, v in (model.get("reference_videos") or {}).items() if v is not None} + reference_audios = {k: v for k, v in (model.get("reference_audios") or {}).items() if v is not None} + if (first_frame is not None or last_frame is not None) and ( + reference_images or reference_videos or reference_audios + ): + raise ValueError( + "First/last frame and reference media are mutually exclusive. Use frames for an " + "image-to-video prompt, or reference media for a reference-to-video prompt." + ) + if reference_audios and not reference_images and not reference_videos: + raise ValueError("Reference audio cannot be used without a reference image or video.") + + for frame in (first_frame, last_frame): + if frame is not None: + validate_image_aspect_ratio(frame, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(frame, min_width=256, min_height=256) + for image in reference_images.values(): + validate_image_aspect_ratio(image, (2, 5), (5, 2), strict=False) # 0.4 to 2.5 + validate_image_dimensions(image, min_width=256, min_height=256) + + total_video_duration = 0.0 + for i, ref_video in enumerate(reference_videos.values(), 1): + try: + ref_fps = float(ref_video.get_frame_rate()) + except Exception: + ref_fps = 0.0 + if ref_fps and not (23.9 <= ref_fps <= 60.5): + raise ValueError(f"Reference video {i} is {ref_fps:.2f} FPS. Supported range is 23.976-60 FPS.") + try: + dur = ref_video.get_duration() + except Exception: + continue + if dur < 1.8: + raise ValueError(f"Reference video {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_video_duration += dur + if total_video_duration > 15.1: + raise ValueError( + f"Total reference video duration is {total_video_duration:.1f}s. Maximum is 15 seconds." + ) + + total_audio_duration = 0.0 + for i, audio in enumerate(reference_audios.values(), 1): + dur = int(audio["waveform"].shape[-1]) / int(audio["sample_rate"]) + if dur < 1.8: + raise ValueError(f"Reference audio {i} is too short: {dur:.1f}s. Minimum duration is 2 seconds.") + total_audio_duration += dur + if total_audio_duration > 15.1: + raise ValueError( + f"Total reference audio duration is {total_audio_duration:.1f}s. Maximum is 15 seconds." + ) + + content: list = [ + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, video, wait_label="Uploading source video"), + ), + role="base_video", + ), + Hailuo03TextContent(text=model["prompt"]), + ] + if first_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, first_frame, max_images=1, wait_label="Uploading first frame" + ) + )[0], + ), + role="first_frame", + ) + ) + if last_frame is not None: + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, last_frame, max_images=1, wait_label="Uploading last frame" + ) + )[0], + ), + role="last_frame", + ) + ) + for i, image in enumerate(reference_images.values(), 1): + content.append( + Hailuo03ImageContent( + image_url=Hailuo03ImageContentUrl( + url=( + await upload_images_to_comfyapi( + cls, image, max_images=1, wait_label=f"Uploading image {i}" + ) + )[0], + ), + role="reference_image", + ) + ) + for i, ref_video in enumerate(reference_videos.values(), 1): + content.append( + Hailuo03VideoContent( + video_url=Hailuo03VideoContentUrl( + url=await upload_video_to_comfyapi(cls, ref_video, wait_label=f"Uploading video {i}"), + ), + ) + ) + for audio in reference_audios.values(): + content.append( + Hailuo03AudioContent( + audio_url=Hailuo03AudioContentUrl( + url=await upload_audio_to_comfyapi( + cls, + audio, + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ), + ) + ) + + response = await sync_op( + cls, + ApiEndpoint(path=HAILUO_03_REGENERATION_ENDPOINT, method="POST"), + response_model=Hailuo03TaskCreationResponse, + data=Hailuo03RegenerationRequest( + model=HAILUO_03_MODELS[model["model"]], + content=content, + resolution=model["resolution"], + aigc_watermark=watermark, + ), + ) + task_result = await poll_op( + cls, + ApiEndpoint(path=f"{HAILUO_03_QUERY_ENDPOINT}/{response.task_id}"), + response_model=Hailuo03TaskQueryResponse, + status_extractor=lambda r: r.task.status, + failed_statuses=HAILUO_03_FAILED_STATUSES, + poll_interval=10, + ) + video_url = task_result.task.content.url if task_result.task.content else None + if not video_url: + raise Exception(f"No video URL in the response: {task_result.model_dump()}") + return IO.NodeOutput(await download_url_to_video_output(video_url)) + + class MinimaxExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[IO.ComfyNode]]: @@ -950,6 +1541,8 @@ class MinimaxExtension(ComfyExtension): MinimaxHailuo03TextToVideoNode, MinimaxHailuo03FirstLastFrameNode, MinimaxHailuo03ReferenceNode, + MinimaxHailuo03ContextIRNode, + MinimaxHailuo03RegenerateNode, ] diff --git a/comfy_api_nodes/nodes_qwen.py b/comfy_api_nodes/nodes_qwen.py new file mode 100644 index 000000000..3b6c5023c --- /dev/null +++ b/comfy_api_nodes/nodes_qwen.py @@ -0,0 +1,442 @@ +import math +import re + +import torch +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension +from comfy_api_nodes.apis.qwen import ( + QwenImageContentItem, + QwenImageGenerationRequest, + QwenImageGenerationResponse, + QwenImageInputField, + QwenImageMessage, + QwenImageParametersField, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + download_url_to_image_tensor, + sync_op, + tensor_to_base64_string, + validate_string, +) + +GENERATION_PATH = "/proxy/qwen/api/v1/services/aigc/multimodal-generation/generation" +QWEN_IMAGE_MODELS = ["qwen-image-3.0-pro", "qwen-image-3.0"] +MIN_AREA = 262144 # 512*512 +MAX_AREA = 6553600 # 2560*2560 +MAX_ASPECT = 8 # the API allows aspect ratios from 1:8 to 8:1 +MAX_INPUT_BYTES = 10 * 1024 * 1024 # the API rejects decoded input images over 10MB + +_IMAGE_REF_RE = re.compile(r"@image(?P\d*)(?!\w)", re.IGNORECASE | re.ASCII) + + +def _resolve_image_refs(prompt: str, total_images: int) -> str: + """Rewrite @Image1-style references (shared partner-node syntax, 1-based; an unnumbered + @image means the first image) into the plain 'Image N' wording the model resolves + natively. A tag counts only at a word boundary or right after a previous tag, so + adjacent tags like '@Image1@Image2' all resolve while addresses like user@image1.com + pass through untouched.""" + parts = [] + pos = 0 + prev_end = -1 + for match in _IMAGE_REF_RE.finditer(prompt): + start = match.start() + if start > 0 and start != prev_end and (prompt[start - 1].isalnum() or prompt[start - 1] == "_"): + continue + idx = int(match.group("idx") or 1) + if not 1 <= idx <= total_images: + raise ValueError( + f"The prompt references @Image{idx}, but only {total_images} reference images " + f"are connected (a batched input counts once per image)." + ) + parts.append(prompt[pos:start]) + parts.append(f"Image {idx}") + pos = match.end() + prev_end = match.end() + parts.append(prompt[pos:]) + return "".join(parts) + + +def _validate_size(width: int, height: int) -> None: + if not MIN_AREA <= width * height <= MAX_AREA: + raise ValueError( + f"Image area must be between {MIN_AREA} (512x512) and {MAX_AREA} (2560x2560) pixels; " + f"got {width}x{height} = {width * height}." + ) + if width > MAX_ASPECT * height or height > MAX_ASPECT * width: + raise ValueError(f"Aspect ratio must be between 1:8 and 8:1; got {width}x{height}.") + + +def _fit_to_size(width: int, height: int) -> tuple[int, int]: + """Scale dimensions into the supported pixel area and 1:8..8:1 aspect range, preserving + the aspect ratio where possible.""" + if width > MAX_ASPECT * height: + height = math.ceil(width / MAX_ASPECT) + elif height > MAX_ASPECT * width: + width = math.ceil(height / MAX_ASPECT) + area = width * height + if area < MIN_AREA: + scale = math.sqrt(MIN_AREA / area) + width, height = math.ceil(width * scale), math.ceil(height * scale) + elif area > MAX_AREA: + scale = math.sqrt(MAX_AREA / area) + width, height = math.floor(width * scale), math.floor(height * scale) + # rounding can push the ratio a hair past the limit; trimming only ever shrinks the area + return min(width, MAX_ASPECT * height), min(height, MAX_ASPECT * width) + + +def _image_data_uri(image: torch.Tensor) -> str: + """PNG data URI of an RGB view of the image, downscaled to <=2048x2048; falls back to + JPEG when the PNG exceeds the API's decoded-size cap (e.g. noisy, incompressible images).""" + image = image[..., :3] + b64 = tensor_to_base64_string(image, total_pixels=2048 * 2048) + if len(b64) * 3 > MAX_INPUT_BYTES * 4: + return "data:image/jpeg;base64," + tensor_to_base64_string( + image, total_pixels=2048 * 2048, mime_type="image/jpeg" + ) + return "data:image/png;base64," + b64 + + +async def _download_result_images(response: QwenImageGenerationResponse) -> torch.Tensor: + if not response.output: + raise Exception(f"An unknown error occurred: {response.code} - {response.message}") + urls = [ + item.image + for choice in response.output.choices + if choice.message + for item in choice.message.content + if item.image + ] + if not urls: + raise Exception(f"The response contains no images: {response.code} - {response.message}") + return torch.cat([await download_url_to_image_tensor(url) for url in urls]) + + +def _size_inputs() -> list[IO.Int.Input]: + return [ + IO.Int.Input( + "width", + default=1024, + min=256, + max=2560, + step=16, + tooltip="The total pixel area must be between 512x512 and 2560x2560; " + "any aspect ratio within that area works.", + ), + IO.Int.Input( + "height", + default=1024, + min=256, + max=2560, + step=16, + tooltip="The total pixel area must be between 512x512 and 2560x2560; " + "any aspect ratio within that area works.", + ), + ] + + +def _t2i_model_option(model_id: str) -> IO.DynamicCombo.Option: + return IO.DynamicCombo.Option( + model_id, + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the image. Supports English and Chinese.", + ), + IO.String.Input( + "negative_prompt", + multiline=True, + default="", + tooltip="Negative prompt describing what to avoid.", + ), + *_size_inputs(), + ], + ) + + +def _edit_model_option(model_id: str) -> IO.DynamicCombo.Option: + return IO.DynamicCombo.Option( + model_id, + [ + IO.Autogrow.Input( + "images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("image"), + names=["image_1", "image_2", "image_3"], + min=1, + ), + tooltip="1-3 reference images. Refer to them in the prompt as @Image1, @Image2, " + "@Image3, numbered in input order; a batched input counts once per image.", + ), + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Editing instructions. Supports English and Chinese, " + "and @Image1-style references to the input images.", + ), + IO.String.Input( + "negative_prompt", + multiline=True, + default="", + tooltip="Negative prompt describing what to avoid.", + ), + ], + ) + + +class QwenImageTextToImageApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="QwenImageTextToImageApi", + display_name="Qwen Image 3 Text to Image", + category="partner/image/Qwen", + description="Generates images from a text prompt using the Qwen-Image 3.0 models.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[_t2i_model_option(model_id) for model_id in QWEN_IMAGE_MODELS], + tooltip="Model to use.", + ), + IO.Int.Input( + "n", + default=1, + min=1, + max=6, + display_mode=IO.NumberDisplay.number, + tooltip="Number of images to generate, returned as a batch.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Image.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.width", "model.height", "n"]), + expr=""" + ( + $isPro := widgets.model = "qwen-image-3.0-pro"; + $area := $lookup(widgets, "model.width") * $lookup(widgets, "model.height"); + $rate := $isPro ? ($area > 2250000 ? 0.10725 : 0.0572) : 0.0429; + {"type":"usd","usd": $rate * widgets.n} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + n: int = 1, + seed: int = 42, + prompt_extend: bool = True, + watermark: bool = False, + ): + validate_string(model["prompt"], strip_whitespace=False, min_length=1) + width, height = model["width"], model["height"] + _validate_size(width, height) + response = await sync_op( + cls, + ApiEndpoint(path=GENERATION_PATH, method="POST"), + response_model=QwenImageGenerationResponse, + data=QwenImageGenerationRequest( + model=model["model"], + input=QwenImageInputField( + messages=[QwenImageMessage(content=[QwenImageContentItem(text=model["prompt"])])], + ), + parameters=QwenImageParametersField( + size=f"{width}*{height}", + n=n, + seed=seed, + prompt_extend=prompt_extend, + watermark=watermark, + negative_prompt=model["negative_prompt"] or None, + ), + ), + ) + return IO.NodeOutput(await _download_result_images(response)) + + +class QwenImageEditApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="QwenImageEditApi", + display_name="Qwen Image 3 Edit", + category="partner/image/Qwen", + description="Edits or combines up to 3 reference images guided by a text prompt " + "using the Qwen-Image 3.0 models.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[_edit_model_option(model_id) for model_id in QWEN_IMAGE_MODELS], + tooltip="Model to use.", + ), + IO.DynamicCombo.Input( + "size", + options=[ + IO.DynamicCombo.Option("match input", []), + IO.DynamicCombo.Option("auto", []), + IO.DynamicCombo.Option("custom", _size_inputs()), + ], + tooltip="Output resolution. 'match input' reuses the first reference image's size, " + "'auto' lets the model pick a size with the same aspect ratio, " + "'custom' sets an explicit width and height.", + ), + IO.Int.Input( + "n", + default=1, + min=1, + max=6, + display_mode=IO.NumberDisplay.number, + tooltip="Number of images to generate, returned as a batch.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Image.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends( + widgets=["model", "size", "size.width", "size.height", "n"], + input_groups=["model.images"], + ), + expr=""" + ( + $isPro := widgets.model = "qwen-image-3.0-pro"; + $mode := widgets.size; + $count := $max([$lookup(inputGroups, "model.images"), 1]); + $inputCost := 0.00429 * $count; + $area := $mode = "custom" + ? $lookup(widgets, "size.width") * $lookup(widgets, "size.height") : 0; + $customRate := $area > 2250000 ? 0.10725 : 0.0572; + $isPro and $mode != "custom" + ? {"type":"range_usd", + "min_usd": 0.0572 * widgets.n + $inputCost, + "max_usd": 0.10725 * widgets.n + $inputCost} + : {"type":"usd", + "usd": ($isPro ? $customRate : 0.0429) * widgets.n + $inputCost} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + size: dict, + n: int = 1, + seed: int = 42, + prompt_extend: bool = True, + watermark: bool = False, + ): + validate_string(model["prompt"], strip_whitespace=False, min_length=1) + reference_images = [image for key in model["images"] for image in model["images"][key]] + if len(reference_images) > 3: + raise ValueError( + f"A maximum of 3 reference images is supported; got {len(reference_images)} " + f"(a batched input counts once per image)." + ) + prompt = _resolve_image_refs(model["prompt"], len(reference_images)) + if size["size"] == "custom": + _validate_size(size["width"], size["height"]) + size_str = f"{size['width']}*{size['height']}" + elif size["size"] == "match input": + height, width = reference_images[0].shape[0], reference_images[0].shape[1] + width, height = _fit_to_size(width, height) + size_str = f"{width}*{height}" + else: # auto: the API picks a size preserving the input aspect ratio (1.9-4.2 MP) + size_str = None + content = [QwenImageContentItem(image=_image_data_uri(image)) for image in reference_images] + content.append(QwenImageContentItem(text=prompt)) + response = await sync_op( + cls, + ApiEndpoint(path=GENERATION_PATH, method="POST"), + response_model=QwenImageGenerationResponse, + data=QwenImageGenerationRequest( + model=model["model"], + input=QwenImageInputField(messages=[QwenImageMessage(content=content)]), + parameters=QwenImageParametersField( + size=size_str, + n=n, + seed=seed, + prompt_extend=prompt_extend, + watermark=watermark, + negative_prompt=model["negative_prompt"] or None, + ), + ), + ) + return IO.NodeOutput(await _download_result_images(response)) + + +class QwenApiExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + QwenImageTextToImageApi, + QwenImageEditApi, + ] + + +async def comfy_entrypoint() -> QwenApiExtension: + return QwenApiExtension() diff --git a/comfy_api_nodes/nodes_recraft.py b/comfy_api_nodes/nodes_recraft.py index 2605b9021..9f1823426 100644 --- a/comfy_api_nodes/nodes_recraft.py +++ b/comfy_api_nodes/nodes_recraft.py @@ -27,6 +27,7 @@ from comfy_api_nodes.util import ( ApiEndpoint, bytesio_to_image_tensor, download_url_as_bytesio, + pad_images_to_common_channels, resize_mask_to_image, sync_op, tensor_to_bytesio, @@ -621,7 +622,7 @@ class RecraftImageToImageNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftImageInpaintingNode(IO.ComfyNode): @@ -723,7 +724,7 @@ class RecraftImageInpaintingNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftTextToVectorNode(IO.ComfyNode): @@ -954,7 +955,7 @@ class RecraftReplaceBackgroundNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftRemoveBackgroundNode(IO.ComfyNode): @@ -995,7 +996,7 @@ class RecraftRemoveBackgroundNode(IO.ComfyNode): image=image[i], path="/proxy/recraft/images/removeBackground", ) - images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) + images.append(torch.cat([bytesio_to_image_tensor(x, mode="RGBA") for x in sub_bytes], dim=0)) pbar.update(1) images_tensor = torch.cat(images, dim=0) @@ -1047,7 +1048,7 @@ class RecraftCrispUpscaleNode(IO.ComfyNode): images.append(torch.cat([bytesio_to_image_tensor(x) for x in sub_bytes], dim=0)) pbar.update(1) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(pad_images_to_common_channels(images), dim=0)) class RecraftCreativeUpscaleNode(RecraftCrispUpscaleNode): diff --git a/comfy_api_nodes/nodes_topaz.py b/comfy_api_nodes/nodes_topaz.py index 9a0c70b4d..d8d051224 100644 --- a/comfy_api_nodes/nodes_topaz.py +++ b/comfy_api_nodes/nodes_topaz.py @@ -217,7 +217,6 @@ class TopazImageEnhance(IO.ComfyNode): response_model=ImageStatusResponse, status_extractor=lambda x: x.status, progress_extractor=lambda x: getattr(x, "progress", 0), - price_extractor=lambda x: x.credits * 0.08, poll_interval=8.0, estimated_duration=60, ) @@ -567,7 +566,6 @@ class TopazImageEnhanceV2(IO.ComfyNode): response_model=ImageStatusResponse, status_extractor=lambda x: x.status, progress_extractor=lambda x: getattr(x, "progress", 0), - price_extractor=lambda x: x.credits * (0.08 if model_choice == "Reimagine" else 0.1144), poll_interval=8.0, estimated_duration=60, ) @@ -814,7 +812,6 @@ class TopazVideoEnhance(IO.ComfyNode): response_model=VideoStatusResponse, status_extractor=lambda x: x.status, progress_extractor=lambda x: getattr(x, "progress", 0), - price_extractor=lambda x: (x.estimates.cost[0] * 0.08 if x.estimates and x.estimates.cost[0] else None), poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_video_output(final_response.download.url)) @@ -1158,7 +1155,6 @@ class TopazVideoEnhanceV2(IO.ComfyNode): response_model=VideoStatusResponse, status_extractor=lambda x: x.status, progress_extractor=lambda x: getattr(x, "progress", 0), - price_extractor=lambda x: (x.estimates.cost[0] * 0.08 if x.estimates and x.estimates.cost[0] else None), poll_interval=10.0, ) return IO.NodeOutput(await download_url_to_video_output(final_response.download.url)) diff --git a/comfy_api_nodes/nodes_tripo.py b/comfy_api_nodes/nodes_tripo.py index 228fe8a1d..10aefa189 100644 --- a/comfy_api_nodes/nodes_tripo.py +++ b/comfy_api_nodes/nodes_tripo.py @@ -66,7 +66,6 @@ async def poll_until_finished( ], status_extractor=lambda x: x.data.status, progress_extractor=lambda x: x.data.progress, - price_extractor=lambda x: x.data.consumed_credit * 0.01 if x.data.consumed_credit else None, estimated_duration=average_duration, ) if response_poll.data.status == TripoTaskStatus.SUCCESS: diff --git a/comfy_api_nodes/nodes_vidu.py b/comfy_api_nodes/nodes_vidu.py index 8c5a43f5b..702417a22 100644 --- a/comfy_api_nodes/nodes_vidu.py +++ b/comfy_api_nodes/nodes_vidu.py @@ -54,7 +54,6 @@ async def execute_task( response_model=TaskStatusResponse, status_extractor=lambda r: r.state, progress_extractor=lambda r: r.progress, - price_extractor=lambda r: r.credits * 0.005 if r.credits is not None else None, max_poll_attempts=max_poll_attempts, ) if not response.creations: diff --git a/comfy_api_nodes/util/__init__.py b/comfy_api_nodes/util/__init__.py index 1fb6b96cf..2bb4a1b04 100644 --- a/comfy_api_nodes/util/__init__.py +++ b/comfy_api_nodes/util/__init__.py @@ -18,6 +18,7 @@ from .conversions import ( downscale_image_tensor_by_max_side, downscale_video_to_max_pixels, image_tensor_pair_to_batch, + pad_images_to_common_channels, pil_to_bytesio, resize_mask_to_image, tensor_to_base64_string, @@ -92,6 +93,7 @@ __all__ = [ "downscale_image_tensor_by_max_side", "downscale_video_to_max_pixels", "image_tensor_pair_to_batch", + "pad_images_to_common_channels", "pil_to_bytesio", "resize_mask_to_image", "tensor_to_base64_string", diff --git a/comfy_api_nodes/util/conversions.py b/comfy_api_nodes/util/conversions.py index f46cac3f8..eb81447a0 100644 --- a/comfy_api_nodes/util/conversions.py +++ b/comfy_api_nodes/util/conversions.py @@ -16,12 +16,14 @@ from comfy_api.latest import Input, InputImpl, Types from ._helpers import mimetype_to_extension -def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor: +def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str | None = None) -> torch.Tensor: """Converts image data from BytesIO to a torch.Tensor. Args: image_bytesio: BytesIO object containing the image data. - mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA"). + mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA"). Defaults + to RGBA when the decoded image carries transparency and RGB when it + does not, so an API that returns no alpha does not get an opaque one. Returns: A torch.Tensor representing the image (1, H, W, C). @@ -31,6 +33,8 @@ def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch ValueError: If the specified mode is invalid. """ image = Image.open(image_bytesio) + if mode is None: + mode = "RGBA" if "A" in image.getbands() or "transparency" in image.info else "RGB" image = image.convert(mode) image_array = np.array(image).astype(np.float32) / 255.0 return torch.from_numpy(image_array).unsqueeze(0) @@ -53,6 +57,17 @@ def image_tensor_pair_to_batch(image1: torch.Tensor, image2: torch.Tensor) -> to return torch.cat((image1, image2), dim=0) +def pad_images_to_common_channels(images: list[torch.Tensor]) -> list[torch.Tensor]: + """Pads [B, H, W, C] image tensors with opaque alpha so they all share the largest channel count.""" + channels = max(image.shape[-1] for image in images) + return [ + torch.nn.functional.pad(image, (0, channels - image.shape[-1]), value=1.0) + if image.shape[-1] < channels + else image + for image in images + ] + + def tensor_to_bytesio( image: torch.Tensor, *, diff --git a/comfy_execution/jobs.py b/comfy_execution/jobs.py index 34c06363b..60f9b8f90 100644 --- a/comfy_execution/jobs.py +++ b/comfy_execution/jobs.py @@ -197,6 +197,7 @@ def normalize_queue_item(item: tuple, status: str) -> dict: 'priority': priority, 'create_time': create_time, 'outputs_count': 0, + 'previewable_outputs_count': 0, 'workflow_id': workflow_id, }) @@ -215,6 +216,7 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs: outputs = history_item.get('outputs', {}) outputs_count, preview_output = get_outputs_summary(outputs) + previewable_outputs_count = count_previewable_outputs(outputs) execution_error = None execution_start_time = None @@ -251,6 +253,7 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs: 'execution_end_time': execution_end_time, 'execution_error': execution_error, 'outputs_count': outputs_count, + 'previewable_outputs_count': previewable_outputs_count, 'preview_output': preview_output, 'workflow_id': workflow_id, }) @@ -345,6 +348,33 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]: return count, preview_output or fallback_preview or text_file_fallback or text_fallback +def count_previewable_outputs(outputs: dict) -> int: + """ + Count only outputs that would actually render in the expanded asset view, + i.e. items is_previewable() accepts (image/video/audio/3D/text). Kept + separate from get_outputs_summary()'s outputs_count, which counts every + output item regardless of media type, so a job with a non-previewable + saved file alongside real media (e.g. SaveLatent's .latent output next to + a SaveImage output) doesn't inflate the Media Assets badge beyond what + the expanded view shows. + """ + count = 0 + for node_outputs in outputs.values(): + if not isinstance(node_outputs, dict): + continue + for media_type, items in node_outputs.items(): + if media_type == 'animated' or not isinstance(items, list): + continue + for item in items: + if not isinstance(item, dict): + item = normalize_output_item(item) + if item is None: + continue + if is_previewable(media_type, item): + count += 1 + return count + + def apply_sorting(jobs: list[dict], sort_by: str, sort_order: str) -> list[dict]: """Sort jobs list by specified field and order.""" reverse = (sort_order == 'desc') diff --git a/comfy_extras/nodes_compositor.py b/comfy_extras/nodes_compositor.py index 10cea3adc..f66a9feab 100644 --- a/comfy_extras/nodes_compositor.py +++ b/comfy_extras/nodes_compositor.py @@ -486,7 +486,10 @@ class ImageCompositor(io.ComfyNode): node_id="ImageCompositor", display_name="Create Layered Image", category="image", + search_aliases=["compositor", "composite", "layer", "layers", "layer editor", "psd"], is_experimental=True, + # both flags on purpose: terminal compositor graphs must execute (the + # editor needs a run to open), and cache hits must replay the layer UI is_output_node=True, has_intermediate_output=True, inputs=[ @@ -605,7 +608,7 @@ class AddLayer(io.ComfyNode): options=list(_LAYER_MODES), default="normal", optional=True, - tooltip="Initial blend mode.", + tooltip="Initial blend mode, applied against the layers below. On the bottom layer over the default transparent background, non-normal modes produce transparency.", ), io.Float.Input( "rotation", diff --git a/comfy_extras/nodes_custom_sampler.py b/comfy_extras/nodes_custom_sampler.py index e81b6328b..c73a8f6dc 100644 --- a/comfy_extras/nodes_custom_sampler.py +++ b/comfy_extras/nodes_custom_sampler.py @@ -591,7 +591,7 @@ class SamplerER_SDE(io.ComfyNode): inputs=[ io.Combo.Input("solver_type", options=["ER-SDE", "Reverse-time SDE", "ODE"]), io.Int.Input("max_stage", default=3, min=1, max=3, advanced=True), - io.Float.Input("eta", default=1.0, min=0.0, max=100.0, step=0.01, round=False, tooltip="Stochastic strength of reverse-time SDE.\nWhen eta=0, it reduces to deterministic ODE. This setting doesn't apply to ER-SDE solver type.", advanced=True), + io.Float.Input("eta", default=1.0, min=0.0, max=10.0, step=0.01, round=False, tooltip="Stochastic strength of SDEs.\nWhen eta=0, they reduce to deterministic ODE.\nLarge eta may cause invalid outputs. If this occurs, try decreasing this value.", advanced=True), io.Float.Input("s_noise", default=1.0, min=0.0, max=100.0, step=0.01, round=False, advanced=True), ], outputs=[io.Sampler.Output()] @@ -599,21 +599,35 @@ class SamplerER_SDE(io.ComfyNode): @classmethod def execute(cls, solver_type, max_stage, eta, s_noise) -> io.NodeOutput: - if solver_type == "ODE" or (solver_type == "Reverse-time SDE" and eta == 0): - eta = 0 - s_noise = 0 + # Extend existing noise scalers phi(x) with eta-controlled noise scalers: + # psi(x) = x**(1-eta) * phi(x)**eta + # where eta is constant and directly scales the h^2(t) contribution. - def reverse_time_sde_noise_scaler(x): + def er_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor: + return x * ((x ** 0.3).exp() + 10.0) ** eta + + def reverse_time_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor: return x ** (eta + 1) - if solver_type == "ER-SDE": - # Use the default one in sample_er_sde() - noise_scaler = None - else: - noise_scaler = reverse_time_sde_noise_scaler + def ode_noise_scaler(x: torch.Tensor) -> torch.Tensor: + return x + + solver_scalers = { + "ER-SDE": er_sde_noise_scaler, + "Reverse-time SDE": reverse_time_sde_noise_scaler, + "ODE": ode_noise_scaler, + } + + if solver_type == "ODE" or eta == 0: + s_noise = 0.0 + solver_type = "ODE" + noise_scaler = solver_scalers[solver_type] sampler_name = "er_sde" - sampler = comfy.samplers.ksampler(sampler_name, {"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage}) + sampler = comfy.samplers.ksampler( + sampler_name, + {"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage}, + ) return io.NodeOutput(sampler) get_sampler = execute @@ -704,15 +718,7 @@ class Noise_EmptyNoise: self.seed = 0 def generate_noise(self, input_latent): - latent_image = input_latent["samples"] - if latent_image.is_nested: - tensors = latent_image.unbind() - zeros = [] - for t in tensors: - zeros.append(torch.zeros(t.shape, dtype=t.dtype, layout=t.layout, device="cpu")) - return comfy.nested_tensor.NestedTensor(zeros) - else: - return torch.zeros(latent_image.shape, dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + return comfy.sample.prepare_empty_noise(input_latent["samples"]) class Noise_RandomNoise: diff --git a/comfy_extras/nodes_dataset.py b/comfy_extras/nodes_dataset.py index 71e5ee368..5ca5a8e14 100644 --- a/comfy_extras/nodes_dataset.py +++ b/comfy_extras/nodes_dataset.py @@ -692,6 +692,7 @@ class ImageProcessingNode(io.ComfyNode): category=cls.category, description=cls.description, is_experimental=True, + is_deprecated=cls.is_deprecated, is_input_list=is_group, # True for group, False for individual inputs=inputs, outputs=[ @@ -861,9 +862,12 @@ class TextProcessingNode(io.ComfyNode): return io.Schema( node_id=cls.node_id, + search_aliases=cls.search_aliases, display_name=cls.display_name or cls.node_id, category="text", + description=cls.description, is_experimental=True, + is_deprecated=cls.is_deprecated, is_input_list=is_group, # True for group, False for individual inputs=inputs, outputs=[ diff --git a/comfy_extras/nodes_lt.py b/comfy_extras/nodes_lt.py index 8c85c92b1..a6e5c5d27 100644 --- a/comfy_extras/nodes_lt.py +++ b/comfy_extras/nodes_lt.py @@ -2,11 +2,14 @@ import nodes import node_helpers import torch import torchaudio +import comfy.ldm.lightricks.duration_head import comfy.model_management import comfy.model_sampling import comfy.samplers import comfy.utils +import logging import math +import re import numpy as np import av from io import BytesIO @@ -934,6 +937,243 @@ class LTXVReferenceAudio(io.ComfyNode): return io.NodeOutput(m, positive, negative) +class LTXVSpatioTemporalGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVSpatioTemporalGuidance", + display_name="LTXV Spatio-Temporal Guidance (STG)", + category="advanced/guidance", + description="Runs one extra pass per step with the self-attention of the selected blocks degraded to a value-passthrough, " + "then guides away from it - improving spatial detail and motion coherence.", + inputs=[ + io.Model.Input("model"), + io.Float.Input("scale", default=1.0, min=0.0, max=100.0, step=0.01, round=0.01), + io.String.Input("blocks", default="29", tooltip="Comma-separated transformer block indices to perturb."), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, scale, blocks, start_percent, end_percent) -> io.NodeOutput: + block_set = frozenset(int(b) for b in re.findall(r"\d+", blocks)) + + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if scale == 0 or not block_set: + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["stg_self_attn_blocks"] = block_set + model_options["transformer_options"] = transformer_options + + (perturbed,) = comfy.samplers.calc_cond_batch(args["model"], [cond], x, args["sigma"], model_options) + + return cfg_result + (cond_pred - perturbed) * scale + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class LTXVModalityGuidance(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVModalityGuidance", + display_name="LTXV Modality Guidance (A/V coupling)", + category="advanced/guidance", + description="Cross-modal (audio-video) guidance for LTXV-AV. Runs one extra forward " + "pass per step with the a2v/v2a cross-attention severed, then pushes the " + "result toward the coupled prediction - strengthening audio-visual sync " + "(e.g. lip-sync). Reference default modality_scale is 3.0. Stacks with the " + "dual-CFG guider and STG. Set to 1.0 to disable (no extra pass).", + inputs=[ + io.Model.Input("model"), + io.Float.Input("modality_scale", default=3.0, min=1.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("start_percent", default=0.0, min=0.0, max=1.0, step=0.001, advanced=True), + io.Float.Input("end_percent", default=1.0, min=0.0, max=1.0, step=0.001, advanced=True), + ], + outputs=[io.Model.Output()], + ) + + @classmethod + def execute(cls, model, modality_scale, start_percent, end_percent) -> io.NodeOutput: + m = model.clone() + model_sampling = m.get_model_object("model_sampling") + sigma_start = model_sampling.percent_to_sigma(start_percent) + sigma_end = model_sampling.percent_to_sigma(end_percent) + + def post_cfg_function(args): + if math.isclose(modality_scale, 1.0): + return args["denoised"] + + sigma_ = args["sigma"][0].item() + if sigma_ > sigma_start or sigma_ < sigma_end: + return args["denoised"] + + cond_pred = args["cond_denoised"] + cond = args["cond"] + cfg_result = args["denoised"] + x = args["input"] + + # Extra pass with audio-video cross-attention severed (both directions) + model_options = args["model_options"].copy() + transformer_options = model_options.get("transformer_options", {}).copy() + transformer_options["a2v_cross_attn"] = False + transformer_options["v2a_cross_attn"] = False + model_options["transformer_options"] = transformer_options + + (mod_pred,) = comfy.samplers.calc_cond_batch( + args["model"], [cond], x, args["sigma"], model_options + ) + + # (modality_scale - 1) * (cond - uncond_modality), per the reference guider. + return cfg_result + (cond_pred - mod_pred) * (modality_scale - 1.0) + + m.set_model_sampler_post_cfg_function(post_cfg_function) + return io.NodeOutput(m) + + +class Guider_LTXAVDualCFG(comfy.samplers.CFGGuider): + """CFG guider that applies separate guidance scales to the video and audio + modalities of a packed LTXV-AV latent. + """ + + def set_conds(self, positive, negative): + self.inner_set_conds({"positive": positive, "negative": negative}) + + def set_cfg(self, video_cfg, audio_cfg): + self.video_cfg = video_cfg + self.audio_cfg = audio_cfg + self.cfg = max(video_cfg, audio_cfg) + + def sample(self, noise, latent_image, *args, **kwargs): + # Capture the video/audio split from the nested latent before it is packed. + self._v_numel = None + if getattr(latent_image, "is_nested", False): + parts = latent_image.unbind() + if len(parts) >= 2: + self._v_numel = math.prod(parts[0].shape[1:]) + return super().sample(noise, latent_image, *args, **kwargs) + + def predict_noise(self, x, timestep, model_options={}, seed=None): + v = getattr(self, "_v_numel", None) + if v is None or math.isclose(self.video_cfg, self.audio_cfg): + # Not an AV latent, or equal scales: fall back to standard single-CFG. + self.cfg = self.video_cfg + return super().predict_noise(x, timestep, model_options, seed) + + video_cfg, audio_cfg = self.video_cfg, self.audio_cfg + + def dual_cfg(args): + # Noise-space: cond = x - cond_pred, uncond = x - uncond_pred; the + # returned tensor is subtracted from x by cfg_function. + cond, uncond = args["cond"], args["uncond"] + out = uncond + (cond - uncond) * video_cfg + out[..., v:] = uncond[..., v:] + (cond[..., v:] - uncond[..., v:]) * audio_cfg + return out + + # disable_cfg1_optimization so the uncond pass always runs even if one of the two scales is 1.0. + model_options = {**model_options, "sampler_cfg_function": dual_cfg, "disable_cfg1_optimization": True} + return super().predict_noise(x, timestep, model_options, seed) + + +class LTXVDualCFGGuider(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDualCFGGuider", + display_name="LTXV Dual CFG Guider", + category="model/sampling/guiders", + description="Separate CFG scales for the video and audio modalities of a packed LTXV-AV latent.", + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Float.Input("video_cfg", default=3.0, min=0.0, max=100.0, step=0.1, round=0.01), + io.Float.Input("audio_cfg", default=7.0, min=0.0, max=100.0, step=0.1, round=0.01), + ], + outputs=[io.Guider.Output()], + ) + + @classmethod + def execute(cls, model, positive, negative, video_cfg, audio_cfg) -> io.NodeOutput: + guider = Guider_LTXAVDualCFG(model) + guider.set_conds(positive, negative) + guider.set_cfg(video_cfg, audio_cfg) + return io.NodeOutput(guider) + + +class LTXVDurationPredictor(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="LTXVDurationPredictor", + display_name="LTXV Duration Predictor", + category="conditioning/video_models", + description="Predicts the natural shot duration for a prompt using the LTX 2.4 duration " + "head (loaded with ModelPatchLoader), and snaps it to the VAE's 8k+1 frame grid.", + search_aliases=["auto duration", "duration head", "num_frames"], + inputs=[ + io.Model.Input("model"), + io.Conditioning.Input("positive"), + io.Custom("MODEL_PATCH").Input("duration_head", + tooltip="LTX 2.4 duration head loaded with ModelPatchLoader."), + io.Float.Input("frame_rate", default=24.0, min=1.0, max=120.0, step=0.01), + io.Float.Input("min_seconds", default=1.0, min=0.5, max=120.0, step=0.1), + io.Float.Input("max_seconds", default=20.0, min=0.5, max=120.0, step=0.1), + ], + outputs=[ + io.Int.Output(display_name="num_frames"), + io.Float.Output(display_name="seconds", tooltip="Raw (unclamped) predicted duration."), + ], + ) + + @classmethod + def execute(cls, model, positive, duration_head, frame_rate, min_seconds, max_seconds) -> io.NodeOutput: + dm = model.model.diffusion_model + head = duration_head.model + if not isinstance(head, comfy.ldm.lightricks.duration_head.DurationHead): + raise ValueError("The connected model_patch is not an LTX duration head.") + + context = positive[0][0] + meta = positive[0][1] + if context.shape[0] != 1: + context = context[:1] + + # Run the caption connectors exactly the way sampling does. + comfy.model_management.load_models_gpu([model, duration_head]) + device = model.load_device + head = head.to(device) + with torch.no_grad(): + context = context.to(device=device, dtype=model.model.get_dtype_inference()) + processed = dm.preprocess_text_embeds(context, unprocessed=meta.get("unprocessed_ltxav_embeds", False)) + video_tokens = processed[..., :dm.cross_attention_dim].float() + audio_tokens = processed[..., dm.cross_attention_dim:].float() + seconds = float(head(video_tokens, audio_tokens)[0]) + + num_frames = comfy.ldm.lightricks.duration_head.seconds_to_num_frames( + seconds, frame_rate, min_seconds, max_seconds) + logging.info("LTXV duration head predicted %.2fs -> %d frames @ %.2f fps", seconds, num_frames, frame_rate) + return io.NodeOutput(num_frames, seconds) + + class LtxvExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[io.ComfyNode]]: @@ -951,6 +1191,10 @@ class LtxvExtension(ComfyExtension): LTXVConcatAVLatent, LTXVSeparateAVLatent, LTXVReferenceAudio, + LTXVDualCFGGuider, + LTXVModalityGuidance, + LTXVSpatioTemporalGuidance, + LTXVDurationPredictor, ] diff --git a/comfy_extras/nodes_lt_audio.py b/comfy_extras/nodes_lt_audio.py index 3ff18d8d4..0924f3e9e 100644 --- a/comfy_extras/nodes_lt_audio.py +++ b/comfy_extras/nodes_lt_audio.py @@ -173,7 +173,7 @@ class LTXAVTextEncoderLoader(io.ComfyNode): node_id="LTXAVTextEncoderLoader", display_name="Load LTXV Audio Text Encoder", category="model/loaders", - description="Recipes:\nltxav: gemma 3 12B", + description="Recipes:\nltxav: gemma 3 12B or matching gemma 4 model", inputs=[ io.Combo.Input( "text_encoder", diff --git a/comfy_extras/nodes_minimax_h3.py b/comfy_extras/nodes_minimax_h3.py index 0b1840e85..0a08f185f 100644 --- a/comfy_extras/nodes_minimax_h3.py +++ b/comfy_extras/nodes_minimax_h3.py @@ -20,6 +20,7 @@ import comfy.model_sampling import comfy.nested_tensor import comfy.utils import node_helpers +from comfy.ldm.minimax.model import FRAME_PER_TOKEN, FRAME_RESCALE from comfy_api.latest import ComfyExtension, io CANVAS_MULTIPLE = 32 @@ -67,6 +68,16 @@ def _resize(image, width, height, crop): return samples.movedim(1, -1) +def _encode_ref_audio(audio_vae, audio): + waveform = audio["waveform"] # [B, C, L] + sr = audio["sample_rate"] + vae_sr = getattr(audio_vae, "audio_sample_rate", 32000) + if sr != vae_sr: + waveform = torchaudio.functional.resample(waveform, sr, vae_sr) + z = audio_vae.encode(waveform[:1].movedim(1, -1)) # [1, 32, 2, T] + return z, z.shape[-1] + + def _empty_av_latent(width, height, length, batch_size=1): frame_count, latent_t, audio_t = temporal_shape(length) video = torch.zeros([batch_size, 24, latent_t, height // 16, width // 16], @@ -144,13 +155,87 @@ class MiniMaxH3ImageToVideo(io.ComfyNode): if keyframes: for kf in keyframes: kf["latent"] = vae.encode(kf.pop("image")) - cond = node_helpers.conditioning_set_values(cond, { - "minimax_keyframes": keyframes, - "minimax_frame_count": frame_count, - }) + cond = node_helpers.conditioning_set_values(cond, {"minimax_keyframes": keyframes}) return io.NodeOutput(cond, latent) +class MiniMaxH3AddGuide(io.ComfyNode): + """Anchor image and/or audio guides at an arbitrary pixel frame of the target video.""" + + @classmethod + def define_schema(cls): + return io.Schema( + node_id="MiniMaxH3AddGuide", + display_name="Add Guide for MiniMax H3", + category="model/conditioning/minimax", + description="Anchor an image, a short clip, audio, or a clip with its soundtrack at any frame of a MiniMax H3 video. Chain several nodes to anchor several frames.", + inputs=[ + io.Conditioning.Input("positive"), + io.Vae.Input("vae", optional=True, tooltip="Video VAE, needed when an image is connected."), + io.Vae.Input("audio_vae", optional=True, tooltip="Audio VAE, needed when an audio is connected."), + io.Latent.Input("latent"), + io.Image.Input("image", optional=True, tooltip="Image or video frames to anchor. Multi-frame batches are anchored as a clip and cropped down to the model's valid clip lengths: 5, 22, 39... (17k + 5) frames. Batches shorter than 5 frames use only the first image."), + io.Audio.Input("audio", optional=True, + tooltip="Soundtrack to anchor starting at the same frame index, cropped to the video's remaining duration."), + io.Int.Input("frame_idx", default=0, min=-9999, max=9999, + tooltip="Frame index to anchor the image or the clip's first frame at. Negative values are counted from the end of the video."), + ], + outputs=[io.Conditioning.Output(display_name="positive")], + ) + + @classmethod + def execute(cls, positive, latent, frame_idx, vae=None, audio_vae=None, image=None, audio=None) -> io.NodeOutput: + samples = latent["samples"] + if not samples.is_nested or len(samples.tensors) != 2 or samples.tensors[0].ndim != 5 or samples.tensors[0].shape[1] != 24: + raise ValueError("MiniMaxH3AddGuide expects a MiniMax H3 AV latent") + if image is None and audio is None: + raise ValueError("MiniMaxH3AddGuide needs an image or an audio to anchor") + video = samples.tensors[0] + height = video.shape[3] * 16 + width = video.shape[4] * 16 + frame_count = sum(FRAME_PER_TOKEN[k % 5] for k in range(video.shape[2])) + + guide_frames = 1 + if image is not None: + if vae is None: + raise ValueError("anchoring guide frames needs the vae input") + guide_frames = image.shape[0] + if guide_frames < 5: + guide_frames = 1 + else: + while guide_frames % 17 != 5: + guide_frames -= 1 + + resolved_frame_index = frame_idx if frame_idx >= 0 else frame_count + frame_idx + if resolved_frame_index < 0 or resolved_frame_index + guide_frames > frame_count: + if guide_frames == 1: + raise ValueError("frame_idx {} is outside the video's {} frames".format(frame_idx, frame_count)) + raise ValueError("a {} frame guide clip at frame_idx {} does not fit in the video's {} frames".format( + guide_frames, frame_idx, frame_count)) + + keyframe = {"resolved_frame_index": resolved_frame_index} + if image is not None: + frames = _resize(image[:guide_frames], width, height, "center") + keyframe["latent"] = vae.encode(frames) + + if audio is not None: + if audio_vae is None: + raise ValueError("anchoring guide audio needs the audio_vae input") + audio_latent, audio_rt = _encode_ref_audio(audio_vae, audio) + # the streams share one time axis: FRAME_RESCALE per pixel frame, 1.0 per audio latent frame + max_rt = math.floor(samples.tensors[1].shape[-1] - FRAME_RESCALE * resolved_frame_index) + if max_rt < 1: + raise ValueError("frame_idx {} is past the end of the video's audio track".format(frame_idx)) + if audio_rt > max_rt: + audio_latent = audio_latent[..., :max_rt].clone() + keyframe["audio_latent"] = audio_latent + + keyframes = list(positive[0][1].get("minimax_keyframes", [])) + keyframes.append(keyframe) + positive = node_helpers.conditioning_set_values(positive, {"minimax_keyframes": keyframes}) + return io.NodeOutput(positive) + + class MiniMaxH3ReferenceToVideo(io.ComfyNode): """ref2va: prompt + reference images / videos / audio -> conditioning + AV latent. @@ -197,16 +282,6 @@ class MiniMaxH3ReferenceToVideo(io.ComfyNode): outputs=[io.Conditioning.Output(display_name="positive"), io.Latent.Output()], ) - @staticmethod - def _encode_ref_audio(audio_vae, audio): - waveform = audio["waveform"] # [B, C, L] - sr = audio["sample_rate"] - vae_sr = getattr(audio_vae, "audio_sample_rate", 32000) - if sr != vae_sr: - waveform = torchaudio.functional.resample(waveform, sr, vae_sr) - z = audio_vae.encode(waveform[:1].movedim(1, -1)) # [1, 32, 2, T] - return z, z.shape[-1] - @classmethod def execute(cls, clip, vae, audio_vae, prompt, width, height, length, ref_image_size="match", ref_images=None, ref_videos=None, ref_video_audios=None, ref_audios=None) -> io.NodeOutput: @@ -254,7 +329,7 @@ class MiniMaxH3ReferenceToVideo(io.ComfyNode): z = vae.encode(frames) audio_latent, ref_audio_t = (None, 0) if soundtrack is not None: - audio_latent, ref_audio_t = cls._encode_ref_audio(audio_vae, soundtrack) + audio_latent, ref_audio_t = _encode_ref_audio(audio_vae, soundtrack) # the soundtrack gets its own