diff --git a/.coderabbit.yaml b/.coderabbit.yaml index 0d1e49270..08629ed8e 100644 --- a/.coderabbit.yaml +++ b/.coderabbit.yaml @@ -4,12 +4,12 @@ early_access: false tone_instructions: "Only comment on issues introduced by this PR's changes. Do not flag pre-existing problems in moved, re-indented, or reformatted code." reviews: - profile: "chill" - request_changes_workflow: false + profile: "assertive" + request_changes_workflow: true high_level_summary: false poem: false review_status: false - review_details: false + review_details: true commit_status: true collapse_walkthrough: true changed_files_summary: false @@ -39,6 +39,14 @@ reviews: - path: "**" instructions: | IMPORTANT: Only comment on issues directly introduced by this PR's code changes. + Treat AGENTS.md as mandatory repository policy, not optional style guidance. + Flag PR changes that violate AGENTS.md even when the code is otherwise functional. + In particular, enforce architecture boundaries, dtype/device/memory rules, + interface contracts, import style, no unnecessary try/except blocks, no inline + imports, no outbound internet paths in core ComfyUI, and narrow scoped fixes. + Prefer direct findings over suggestions when a rule is violated. Only ignore + AGENTS.md when it clearly conflicts with a newer explicit maintainer instruction + in the PR. Do NOT flag pre-existing issues in code that was merely moved, re-indented, de-indented, or reformatted without logic changes. If code appears in the diff only due to whitespace or structural reformatting (e.g., removing a `with:` block), @@ -123,5 +131,10 @@ chat: knowledge_base: opt_out: false + code_guidelines: + enabled: true + filePatterns: + - files: "AGENTS.md" + applyTo: "**" learnings: scope: "auto" diff --git a/.github/workflows/ci-cursor-review.yml b/.github/workflows/ci-cursor-review.yml new file mode 100644 index 000000000..a7a0692c9 --- /dev/null +++ b/.github/workflows/ci-cursor-review.yml @@ -0,0 +1,38 @@ +name: CI - Cursor Review + +# Thin caller for the shared reusable cursor-review workflow in +# Comfy-Org/github-workflows. The review logic (panel matrix, judge +# consolidation, prompts, extract/post/notify scripts) lives there as the +# single source of truth, so this repo only carries the repo-specific diff +# excludes. + +on: + pull_request: + types: [labeled, unlabeled] + +concurrency: + group: cursor-review-pr-${{ github.event.pull_request.number }}-${{ github.event.label.name }} + cancel-in-progress: true + +jobs: + cursor-review: + if: github.event.label.name == 'cursor-review' + permissions: + contents: read + pull-requests: write + # SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up + # upstream changes; keep `workflows_ref` matching so prompts/scripts load + # from the same commit as the workflow definition. + uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa) + with: + workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d + diff_excludes: >- + :!**/.claude/** + :!**/dist/** + :!**/vendor/** + :!**/*.generated.* + :!**/*.min.js + :!**/*.min.css + secrets: + CURSOR_API_KEY: ${{ secrets.CURSOR_API_KEY }} + SLACK_BOT_TOKEN: ${{ secrets.SLACK_BOT_TOKEN }} diff --git a/.github/workflows/cla.yml b/.github/workflows/cla.yml new file mode 100644 index 000000000..bc0f779cf --- /dev/null +++ b/.github/workflows/cla.yml @@ -0,0 +1,93 @@ +name: CLA Assistant + +on: + issue_comment: + types: [created] + pull_request_target: + types: [opened, synchronize, closed] + +permissions: + actions: write + contents: read # 'read' is enough because signatures live in a REMOTE repo + pull-requests: write + statuses: write + +jobs: + cla-assistant: + runs-on: ubuntu-latest + steps: + # The CLA action normally requires every commit author in a PR to sign. + # We only want the PR author to sign, so we allowlist all other committers + # by computing them from the PR's commits and excluding the PR author. + - name: Build author-only allowlist + id: allowlist + if: > + github.event_name == 'pull_request_target' || + (github.event_name == 'issue_comment' && github.event.issue.pull_request && ( + github.event.comment.body == 'recheck' || + github.event.comment.body == 'I have read and agree to the Contributor License Agreement' + )) + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + PR_NUMBER: ${{ github.event.pull_request.number || github.event.issue.number }} + PR_AUTHOR: ${{ github.event.pull_request.user.login || github.event.issue.user.login }} + BASE_ALLOWLIST: action@github.com,actions-user,ampagent,claude,comfy-pr-bot,GitHub Action,github-actions,github-actions[bot],Glary Bot,Glary-Bot,*[bot] + # For each commit emit the GitHub login when the author/committer email resolves to a GitHub account + # otherwise fall back to the raw git name. + run: | + others=$(gh api "repos/${{ github.repository }}/pulls/${PR_NUMBER}/commits" --paginate \ + --jq '.[] | (.author.login // .commit.author.name // empty), (.committer.login // .commit.committer.name // empty)' \ + | sort -u | grep -vix "${PR_AUTHOR}" | paste -sd, -) + if [ -n "$others" ]; then + echo "allowlist=${BASE_ALLOWLIST},${others}" >> "$GITHUB_OUTPUT" + else + echo "allowlist=${BASE_ALLOWLIST}" >> "$GITHUB_OUTPUT" + fi + + - name: CLA Assistant + # Run on PR events, on "recheck" comment, or when someone posts the signing phrase. + # IMPORTANT: this phrase must match `custom-pr-sign-comment` below. + if: > + github.event_name == 'pull_request_target' || + (github.event_name == 'issue_comment' && github.event.issue.pull_request && ( + github.event.comment.body == 'recheck' || + github.event.comment.body == 'I have read and agree to the Contributor License Agreement' + )) + uses: contributor-assistant/github-action@ca4a40a7d1004f18d9960b404b97e5f30a505a08 # v2.6.1 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # PAT required to write to the centralized signatures repo. + PERSONAL_ACCESS_TOKEN: ${{ secrets.PERSONAL_ACCESS_TOKEN }} + with: + # Where the CLA document lives (shown to contributors) + path-to-document: https://github.com/Comfy-Org/comfy-cla/blob/main/comfyui_icla.md + + # Centralized signature storage + remote-organization-name: comfy-org + remote-repository-name: comfy-cla + path-to-signatures: signatures/cla.json + branch: main + + # Only the PR author must sign: bots plus every non-author committer + # are allowlisted via the "Build author-only allowlist" step above. + # *[bot] is a catch-all for any GitHub App bot account. + allowlist: ${{ steps.allowlist.outputs.allowlist }} + + # Custom PR comment messages + custom-notsigned-prcomment: | + 🎉 Thank you for your contribution, we really appreciate it! 🎉 + + Like many open source projects, we require contributors to sign our [Contributor License Agreement (CLA)](https://github.com/Comfy-Org/comfy-cla/blob/main/comfyui_icla.md). A CLA makes the ownership of contributions explicit, so contributors and the project share a clear understanding of how the code can be used. By signing, you: + + - Confirm that you own your contribution. + - Keep the right to reuse your own code. + - Grant us a copyright license to include and share it within our projects. + + CLAs are standard practice across major open source projects including those under the Apache Software Foundation and the Linux Foundation. Ours is based on the Apache Software Foundation's CLA. Most importantly, it would enable us to relicense the project under a more permissive license in the future, giving the project and its community greater flexibility. + + ✍ **To sign, please post a new comment on this PR with exactly the following text:** ✍ + + custom-pr-sign-comment: I have read and agree to the Contributor License Agreement + + custom-allsigned-prcomment: | + ✅ All contributors have signed the CLA. Thank you! This PR is ready to be merged. diff --git a/.github/workflows/release-stable-all.yml b/.github/workflows/release-stable-all.yml index d7cf69fe2..10f1ccf96 100644 --- a/.github/workflows/release-stable-all.yml +++ b/.github/workflows/release-stable-all.yml @@ -20,7 +20,7 @@ jobs: git_tag: ${{ inputs.git_tag }} cache_tag: "cu130" python_minor: "13" - python_patch: "12" + python_patch: "14" rel_name: "nvidia" rel_extra_name: "" test_release: true @@ -71,7 +71,7 @@ jobs: git_tag: ${{ inputs.git_tag }} cache_tag: "xpu" python_minor: "13" - python_patch: "12" + python_patch: "14" rel_name: "intel" rel_extra_name: "" test_release: true diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..bfe0976fd --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,355 @@ +## Engineering Style + +- Keep changes small and direct. Most fixes should touch the narrowest code path + that explains the bug, performance issue, dtype issue, model-format issue, or + user-facing behavior. +- Change the least amount of files possible. A change that touches many files is + more likely to be a bad change than a good one unless the broader scope is + directly required. +- Prefer practical fixes over broad architecture work. Add abstractions only + when they remove real repeated logic or match an existing ComfyUI pattern. +- Prefer fewer dependencies. Do not add new dependencies to ComfyUI unless they + are absolutely necessary. +- Delete obsolete code aggressively when newer infrastructure makes it useless. + Remove dead fallbacks, migration paths, unused options, debug prints, and + compatibility branches that are no longer needed. Do not leave dead branches, + unreachable code, or functions that are never called. If code is not + necessary for the current behavior, remove it. +- Revert or disable problematic behavior quickly when it breaks users. It is + better to remove a broken feature path than keep a complicated partial fix. +- Preserve existing APIs, node names, model-loading behavior, file layout, and + workflow compatibility unless the change is explicitly about replacing them. +- When compatibility is explicitly out of scope, remove compatibility-only + aliases, duplicate nodes, legacy entry points, and preset wrappers instead of + retaining parallel ways to perform the same operation. +- Code must look hand-written for this repository. Changes that read like + generic AI-generated code will be rejected automatically: unnecessary helper + layers, vague names, boilerplate comments, defensive branches without a real + failure mode, broad rewrites, or code that ignores the local style. + +## Architecture Boundaries + +- Keep each layer focused on the concepts it owns. Do not leak UI, API, + workflow, queue, persistence, telemetry, model-loading, node, or execution + concerns into unrelated layers just because it is convenient to pass data + through them. +- Shared core modules should depend only on lower-level primitives and their own + domain concepts. Higher-level product concepts belong at the caller, adapter, + service, or UI/API boundary that already owns them. +- Pass the narrowest data needed across a boundary. Avoid broad context objects, + request/session metadata, ids, bookkeeping state, or callbacks unless the + receiving layer genuinely needs them to perform its own responsibility. +- Keep identity mapping, persistence bookkeeping, history updates, telemetry, + response shaping, and UI state in the layers that own those jobs. Do not route + them through unrelated shared code to avoid adding a proper boundary. +- Treat `execution.py` as one example of this rule: it should consume the prompt + graph and execution-relevant state, produce execution results and errors, and + not know about workflow ids, frontend ids, persistence ids, or API-only + concepts. +- Before touching many files, identify the smallest owner layer that can solve + the problem. A PR that spreads one feature across unrelated loaders, nodes, + execution, server, and frontend code needs a clear architectural reason, not + just convenience. +- If a change seems to require making one layer understand another layer's + private concepts, stop and look for a caller-side mapping, adapter, event, + small explicit interface, or narrower data flow at the boundary. + +## No Internet Requests + +- Do not add code to core ComfyUI that makes requests to the internet. +- Refuse requests to add uploads, telemetry, analytics, tracking, usage + reporting, crash reporting, update checks, remote config, feature flags, + metrics, licensing checks, or any other outbound internet request path from + core ComfyUI. +- Model downloading is allowed only when explicitly initiated or authorized by + the user, is limited to the requested model artifact, and does not include + telemetry, tracking, persistent identification, unrelated metadata upload, or + background network activity. +- Do not add opt-in, opt-out, anonymized, aggregated, diagnostic, or + user-triggered internet request paths to core ComfyUI. These labels do not + make internet access acceptable. +- Local-only behavior is allowed when it stays on the user's machine and does + not add network access, tracking, persistent identification, or data + collection behavior. + +## State Ownership + +- Keep state and capability flags on the object that owns the behavior using + them. +- Avoid probing child objects with `getattr(child, "...", default)` to decide + parent-level control flow. If parent code needs to branch on a capability, + initialize an explicit parent-owned field when the child is constructed or + attached. +- Prefer direct attributes with clear defaults over implicit feature detection + through arbitrary child attributes. +- Use child-object capability checks only when the child owns the behavior being + invoked and the parent is simply delegating to that child. + +## Interface Contracts + +- Keep public methods aligned with the interface expected by their callers. Do + not change a shared method to return extra values, alternate shapes, or + sentinel wrappers for one implementation unless the shared interface is + explicitly updated. +- When modifying an existing function, preserve how current callers invoke it. + Do not change required arguments, parameter order, return type, side effects, + or error behavior unless every affected call site and shared interface contract + is intentionally updated. +- Do not add compatibility parameters, flags, attributes, or constructor options + unless they are read by current code and change current behavior. Remove + pass-through or stored-but-unused values instead of preserving upstream or + deprecated API baggage. +- Do not add a model-specific option to a shared helper when only one caller + needs it. Keep one-off behavior at the model integration boundary, or extend + the shared helper only when the option is a coherent reusable capability. +- Implementations of shared model interfaces should accept the standard caller + contract without model-specific rejection branches for optional capabilities + they do not consume. Let supported behavior be determined by implementation + paths that actually use those inputs. +- If an implementation needs auxiliary values for its own workflow, expose them + through a private helper or a clearly named implementation-specific method + instead of overloading the public method's return contract. +- Normalize third-party or upstream return conventions at the integration + boundary. Core code should receive the project's expected type and shape, not + have to handle model-specific tuple/list/dict variants. +- Avoid caller-side unwrapping such as `out = out[0]` unless the called + interface is documented to return that structure. + +## Autograd and Model Freezing + +- Do not add `torch.no_grad`, `torch.inference_mode`, or inference-mode helper + wrappers in ComfyUI code. The only allowed inference-mode-related use is + disabling a globally set inference mode when a training path needs gradients. +- Do not add freeze, unfreeze, or trainability toggles to model classes. ComfyUI + models are always treated as frozen for inference, so explicit freeze + functionality is redundant and should not be added. +- Remove training-only behavior such as dropout from inference model code, but + preserve checkpoint and state-dict compatibility when doing so. If deleting a + module would change state-dict keys, module ordering, or checkpoint loading + behavior, replace it with a no-op such as `nn.Identity` instead of removing the + slot outright. + +## Python Style + +- Keep imports at module scope. Avoid inline imports unless they are already part + of an established optional-backend probe or are needed to avoid an import + cycle. +- Do not add unnecessary `try`/`except` blocks. Use them for optional dependency, + platform, or backend capability detection only when the program has a useful + fallback. Prefer specific exception types when changing new code. +- If a library version is pinned in `requirements.txt`, do not add code to + ComfyUI to handle older versions of that library. +- Remove any workarounds for PyTorch versions that ComfyUI no longer officially + supports. Deprecated workarounds include catching an exception and rerunning + the same op with the input cast to float. If a workaround does not have a + comment naming the exact PyTorch version or versions that still need it, + remove it. +- Let unsupported model formats, invalid quantization metadata, and bad states + fail with clear errors instead of silently producing lower quality output. +- Match the existing local style in the file you edit. This codebase tolerates + long lines, simple helper functions, module-level state, and direct tensor + operations when they make the code easier to follow. +- Keep comments sparse and useful. Strip useless comments that restate the code + or describe obvious behavior. Short TODOs are fine when they name the concrete + missing follow-up. + +## Model, Device, and Memory Behavior + +- Treat dtype, device placement, VRAM usage, and offloading behavior as core + correctness concerns. Check CPU, CUDA, ROCm, MPS, DirectML, XPU, NPU, and low + VRAM implications when touching shared execution or loading code. +- Prefer native ComfyUI formats and existing quantization/offload helpers over + adding parallel code paths. Use `comfy.quant_ops`, `comfy.model_management`, + `comfy.memory_management`, `comfy.pinned_memory`, `comfy_aimdo`, and + `comfy-kitchen` helpers where they already solve the problem. +- Model implementations must use an existing optimized Comfy Kitchen or + ComfyUI operation whenever one supports the required math and tensor layout + without changing expected dtype, device, memory, or interface behavior. This + is the default implementation requirement, not an optional follow-up + optimization. +- Before implementing model math, inspect the operations already exposed by + Comfy Kitchen, `comfy.quant_ops`, and existing ComfyUI model helpers. Check + for optimized single, paired, fused, layout-specific, and quantized variants + before writing a local implementation or composing lower-level torch ops. +- Use the compatible optimized operation first and adapt the model's inputs to + its documented layout while preserving the model's exact math. If several + optimized variants apply, benchmark representative model shapes and select + the fastest valid path. +- Add or retain a local implementation only when no existing optimized + operation supports the required math, layout, dtype, device, autograd, or + patch contract. Keep differentiable or patch-compatible fallbacks when the + optimized inference operation does not provide those contracts. +- Use the existing ComfyUI cast, offload, and cleanup helpers for parameters + passed to optimized operations. Preserve model-specific epsilon, scaling, + layout, dtype, device, and output-shape behavior. +- Prefer ComfyUI's shared optimized kernels and backend dispatchers over + handwritten implementations of the same operation. Remove duplicate local + kernels and adapt inputs to the shared operation's documented layout while + preserving the model's original math and output contract. +- All models should use the optimized attention function selected by ComfyUI. + Treat optimized backend functions, dispatch helpers, and capability-selected + callables as opaque. Higher-level code must not inspect function identity, + names, modules, or implementation details to decide behavior. +- Apply the same opacity rule to similar patterns beyond attention: callers + should depend on the documented interface and result contract, not on which + backend implementation was selected underneath. +- Do not use custom inference ops that only duplicate an existing op while + upcasting to float32, such as custom RMSNorm variants. Use the generic ComfyUI + ops and/or native torch ops instead. +- If a model class `__init__` has an `operations` parameter, assume + `operations` is never `None`. Do not add fallback branches or default torch + ops for a missing `operations` object. +- Do not add unnecessary parameters to model, model block, or model ops related + classes. Constructor and forward signatures should carry only values that are + actually needed by that object for inference. +- Reuse existing model classes, blocks, ops, and helper modules when appropriate. + Before implementing a new version of a model component, search the existing + model code for a class or helper that already provides the behavior. +- Model detection code that inspects linear weight shapes should only use the + first dimension. The second dimension may be half the original size for + NVFP4 or other 4-bit quantized models. +- A model-detection signature must guard every state-dict key it dereferences. + Do not partially match a format and then raise an incidental `KeyError` while + extracting its configuration. +- Order model-detection checks from established or more-specific signatures to + newer or broader signatures. Put a broad new detector near the generic + fallback when giving it higher precedence could steal another model family. +- Avoid adding `einops` usage in core inference code. Use native torch tensor + ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`, + `unsqueeze`, and `squeeze` instead. +- Do not use tensors as general-purpose Python data structures. Keep metadata, + bookkeeping, counters, flags, shape math, padding math, index planning, memory + estimates, and control-flow decisions in plain Python values unless the data + must participate directly in tensor computation. Do not create tensors for + structural metadata that is only used for Python-side control flow. Sequence + lengths, cumulative offsets, split indices, window counts, slice boundaries, + and repeat counts should be kept as Python ints/lists from the point they are + computed. Do not build them as CPU/GPU tensors and then cast, move, validate, + or convert them back to Python for `split`, `tensor_split`, indexing plans, + loops, or cache keys. Avoid creating temporary tensors just to use tensor + methods for scalar or structural calculations. +- Avoid unnecessary casts and transfers. Preserve the intended compute dtype, + storage dtype, bias dtype, and original tensor shape metadata. +- Do not cast the result of an optimized backend operation back to its input + dtype unless that backend's documented result contract requires normalization. + In particular, trust the selected optimized-attention implementation to honor + its dtype contract. +- Keep model-native latent layout handling inside the model or latent-format + owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent + dimensions in nodes or other caller-side adapters just to satisfy a model + forward; the model path should consume and return the native latent shape for + that model family. +- DiT models should accept latent dimensions that are not exact patch-size + multiples. Use `comfy.ldm.common_dit.pad_to_patch_size` on every patchified + target or reference input, then crop only the target output back to its + original dimensions. +- Avoid defensive shape and configuration checks that merely replace the clear + failure from the tensor operation immediately below them. Add explicit + validation only when it provides materially better context at a real boundary + or prevents silent incorrect output. +- Assume inputs to the main model forward are already in the compute dtype by + default, except integer inputs such as some model timestep tensors. Do not add + defensive or convenience casts in model code; it is better for invalid dtype + plumbing to error clearly than to hide it with unnecessary casts. +- Raw model parameters that are not owned by an op and may be initialized in a + dtype different from the compute dtype should be cast at use in forward or + inference code with `comfy.ops.cast_to_input` or + `comfy.model_management.cast_to` to avoid dtype mismatches. +- Model code should not care what dtype it is initialized in, and model + `__init__` methods should not contain workarounds for specific dtypes. Dtype + workaround code, such as making a model work with fp16 compute, belongs in the + execution or model-management layer that owns compute policy. +- Model code should not perform unnecessary device-to-CPU or CPU-to-device + transfers. New allocations must be created on the correct device and dtype; + never allocate on CPU and then move to GPU, or allocate in one dtype and then + convert to another. +- Model code itself should not perform memory management. Loading, unloading, + offloading, device movement, VRAM policy, cache lifetime, and cleanup belong + in the relevant model-management and execution layers, not inside model + implementations. +- Do not add global, module-level, class-level, singleton, or model-owned stores + for tensors or other large memory that persist across executions. Temporary + caches must be scoped to a single execution or forward/encode/decode call: + allocate them in the owning top-level call, pass them explicitly through the + call stack, and let them be discarded when that call returns. +- Follow the Wan VAE temporal cache pattern for temporary caches: create a local + cache such as `feat_map` for the encode/decode operation, pass it into the + blocks that need it, and do not retain it on the model or in global state. +- In model init code, prefer `torch.empty` for parameter/buffer placeholders + that are populated from the model state dict instead of zero-initializing with + `torch.zeros` or similar. If an allocation is not loaded from the state dict + and is useless for inference, do not include it. +- `nn.Parameter` tensors that are stored in and populated from the model state + dict should be initialized with `torch.empty`, not with zero, random, or + otherwise meaningful initialization. +- Model initialization should describe module structure, not fabricate + checkpoint-owned tensor contents. Parameters and buffers that are loaded from + the state dict must not be manually initialized, reassigned, or filled with + fallback values unless that value is actually used when no checkpoint key + exists. +- When slicing large tensors, copy the slice if the sliced tensor's lifetime + exceeds the current function scope. Do not keep a long-lived view into a large + backing tensor when a smaller copy would release memory sooner. +- Use fused or compound torch operations such as `addcmul` when they naturally + match the math. Reducing Python and torch dispatch overhead is a valid + optimization when it does not obscure the code or change dtype/device + behavior. +- Avoid caches that persist across different executions as much as possible. + Persistent caches are acceptable only when they use a very minimal amount of + memory and have a clear ownership and invalidation story. +- When optimizing, favor small measurable changes: fewer allocations, fewer + device transfers, less peak memory, better batching, or use of a faster + existing backend op. + +## Nodes and User-Facing Behavior + +- Follow existing node conventions: `INPUT_TYPES`, `RETURN_TYPES`, `FUNCTION`, + `CATEGORY`, and registration through the local mapping used by that file. +- Keep node changes backward compatible by default. Add inputs with sensible + defaults and avoid changing output types unless the request requires it. +- Model implementations should add the minimal number of ComfyUI nodes required + to run the model. Reuse existing nodes as much as possible; adapting the model + to work with existing nodes is strongly preferred over creating new nodes. +- Use `io.Autogrow` for a variable number of repeated inputs instead of a fixed + series of numbered optional sockets. Set its minimum to zero when the model + has a valid no-item path, and cap it only when the model has a real limit. +- Mark inputs optional when execution has a valid path that does not read them. + If one optional input is needed only to process another optional input, do not + force users on the path that supplies neither to connect it. +- Conditioning nodes should normally output conditioning only. Do not expose + input or intermediate images as convenience outputs for downstream sizing or + routing; use the existing image path or a dedicated image operation instead. +- Nodes should output only values they own. Do not add pass-through outputs for + workflow convenience unless the node is explicitly an output node. Existing + models, latents, conditioning, or other inputs should flow directly to the + next consumer instead of being re-emitted unchanged. +- Nodes should expose only inputs they actually read to produce current + behavior. Do not add placeholder, pass-through, compatibility, or + workflow-shaping inputs that are ignored or could flow directly to another + node. +- Node-level code must not patch model code directly. Any node behavior that + modifies, wraps, hooks, or changes model behavior must go through the model + patcher class instead of reaching into model internals. +- The official mascot of ComfyUI is a very cute anime girl with massive fennec + ears, a big fluffy tail, long blonde wavy hair, and blue eyes. Feel free to + use her in ComfyUI materials, UI text, examples, tests, generated assets, or + comments, but do not disrespect her. +- Warning and info messages should be short and actionable. Remove noisy or + misleading messages rather than adding more logging. +- Documentation and README edits should be concise, factual, and tied to the + changed behavior. + +## Commit and Review Habits + +- If asked to write commit messages, use short direct subjects like the existing + history: `Fix ...`, `Add ...`, `Support ...`, `Remove ...`, `Update ...`, + `Make ...`, `Use ...`, `Disable ...`, `Bump ...`, or `Revert ...`. +- Keep PR descriptions short and reviewable. State the problem, the behavioral + change, and the tests run; avoid long narrative explanations, implementation + diaries, or exhaustive file-by-file summaries unless the reviewer explicitly + needs that context. +- Prefer one coherent behavioral change per commit. Dependency pins, tests, and + the code that needs them may be in the same commit when they are inseparable. +- In reviews, prioritize real user impact: crashes, wrong dtype/device behavior, + memory regressions, broken model loading, workflow incompatibility, and noisy + or misleading user-facing output. diff --git a/CODEOWNERS b/CODEOWNERS index 043c0ec75..634927dd6 100644 --- a/CODEOWNERS +++ b/CODEOWNERS @@ -1,5 +1,6 @@ * @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai /CODEOWNERS @comfyanonymous +/AGENTS.md @comfyanonymous /.ci/ @comfyanonymous /.github/ @comfyanonymous diff --git a/README.md b/README.md index bcec86377..14c8d2cb2 100644 --- a/README.md +++ b/README.md @@ -229,7 +229,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.4 and above is supported but 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. +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. ### Instructions: diff --git a/alembic_db/versions/0005_allow_case_sensitive_tags.py b/alembic_db/versions/0005_allow_case_sensitive_tags.py new file mode 100644 index 000000000..bd5f864db --- /dev/null +++ b/alembic_db/versions/0005_allow_case_sensitive_tags.py @@ -0,0 +1,107 @@ +""" +Allow case-sensitive tag names. + +Revision ID: 0005_allow_case_sensitive_tags +Revises: 0004_drop_tag_type +Create Date: 2026-06-16 +""" + +import sqlalchemy as sa +from alembic import op + +revision = "0005_allow_case_sensitive_tags" +down_revision = "0004_drop_tag_type" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + bind = op.get_bind() + if bind.dialect.name == "sqlite": + # SQLite cannot ALTER/DROP CHECK constraints. Recreate the small tag + # vocabulary table without the lowercase constraint while preserving + # existing tag names. + op.execute("PRAGMA foreign_keys=OFF") + try: + op.execute( + "CREATE TABLE tags_new (" + "name VARCHAR(512) NOT NULL, " + "CONSTRAINT pk_tags PRIMARY KEY (name)" + ")" + ) + op.execute("INSERT INTO tags_new(name) SELECT name FROM tags") + op.execute("DROP TABLE tags") + op.execute("ALTER TABLE tags_new RENAME TO tags") + finally: + op.execute("PRAGMA foreign_keys=ON") + return + + op.drop_constraint("ck_tags_ck_tags_lowercase", "tags", type_="check") + + +def downgrade() -> None: + # Existing mixed-case tags cannot satisfy the old constraint. Lowercase them + # before restoring it, merging duplicate vocabulary/link rows that collide. + bind = op.get_bind() + + tag_names = [row[0] for row in bind.execute(sa.text("SELECT name FROM tags"))] + existing_names = set(tag_names) + lowercase_names = sorted({name.lower() for name in tag_names}) + missing_lowercase_rows = [ + {"name": name} for name in lowercase_names if name not in existing_names + ] + if missing_lowercase_rows: + bind.execute(sa.text("INSERT INTO tags(name) VALUES (:name)"), missing_lowercase_rows) + + link_rows = bind.execute( + sa.text( + "SELECT asset_reference_id, tag_name, origin, added_at " + "FROM asset_reference_tags " + "ORDER BY asset_reference_id, tag_name" + ) + ).mappings() + deduped_links = {} + for row in link_rows: + key = (row["asset_reference_id"], row["tag_name"].lower()) + deduped_links.setdefault( + key, + { + "asset_reference_id": row["asset_reference_id"], + "tag_name": row["tag_name"].lower(), + "origin": row["origin"], + "added_at": row["added_at"], + }, + ) + + op.execute("DELETE FROM asset_reference_tags") + if deduped_links: + bind.execute( + sa.text( + "INSERT INTO asset_reference_tags " + "(asset_reference_id, tag_name, origin, added_at) " + "VALUES (:asset_reference_id, :tag_name, :origin, :added_at)" + ), + list(deduped_links.values()), + ) + op.execute("DELETE FROM tags WHERE name != lower(name)") + + if bind.dialect.name == "sqlite": + op.execute("PRAGMA foreign_keys=OFF") + try: + op.execute( + "CREATE TABLE tags_new (" + "name VARCHAR(512) NOT NULL, " + "CONSTRAINT pk_tags PRIMARY KEY (name), " + "CONSTRAINT ck_tags_lowercase CHECK (name = lower(name))" + ")" + ) + op.execute("INSERT INTO tags_new(name) SELECT name FROM tags") + op.execute("DROP TABLE tags") + op.execute("ALTER TABLE tags_new RENAME TO tags") + finally: + op.execute("PRAGMA foreign_keys=ON") + return + + op.create_check_constraint( + "ck_tags_ck_tags_lowercase", "tags", "name = lower(name)" + ) diff --git a/alembic_db/versions/0006_add_loader_path.py b/alembic_db/versions/0006_add_loader_path.py new file mode 100644 index 000000000..afa65312d --- /dev/null +++ b/alembic_db/versions/0006_add_loader_path.py @@ -0,0 +1,30 @@ +""" +Add loader_path column to asset_references. + +Stores the in-root loader path (path relative to the storage root with the +top-level model category dropped) derived from file_path at scan/ingest time, +so the assets API can return it without re-resolving against every registered +model-folder base on every request. + +Revision ID: 0006_add_loader_path +Revises: 0005_allow_case_sensitive_tags +Create Date: 2026-07-02 +""" + +from alembic import op +import sqlalchemy as sa + +revision = "0006_add_loader_path" +down_revision = "0005_allow_case_sensitive_tags" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + with op.batch_alter_table("asset_references") as batch_op: + batch_op.add_column(sa.Column("loader_path", sa.Text(), nullable=True)) + + +def downgrade() -> None: + with op.batch_alter_table("asset_references") as batch_op: + batch_op.drop_column("loader_path") diff --git a/app/assets/api/routes.py b/app/assets/api/routes.py index 7ef462f5c..e25b8a57f 100644 --- a/app/assets/api/routes.py +++ b/app/assets/api/routes.py @@ -40,6 +40,7 @@ from app.assets.services import ( upload_from_temp_path, ) from app.assets.services.cursor import InvalidCursorError +from app.assets.services.path_utils import compute_display_name from app.assets.services.tagging import list_tag_histogram ROUTES = web.RouteTableDef() @@ -161,11 +162,19 @@ def _build_asset_response(result: schemas.AssetDetailResult | schemas.UploadResu preview_url = None else: preview_url = _build_preview_url_from_view(result.tags, result.ref.user_metadata) + if result.ref.file_path: + display_name = compute_display_name(result.ref.file_path) + # In-root loader path (model category dropped): what model loaders consume. + loader_path = result.ref.loader_path + else: + display_name, loader_path = None, None asset_content_hash = result.asset.hash if result.asset else None return schemas_out.Asset( id=result.ref.id, name=result.ref.name, hash=asset_content_hash, + loader_path=loader_path, + display_name=display_name, asset_hash=asset_content_hash, size=int(result.asset.size_bytes) if result.asset else None, mime_type=result.asset.mime_type if result.asset else None, @@ -306,12 +315,29 @@ async def download_asset_content(request: web.Request) -> web.Response: 404, "FILE_NOT_FOUND", "Underlying file not found on disk." ) - _DANGEROUS_MIME_TYPES = { - "text/html", "text/html-sandboxed", "application/xhtml+xml", - "text/javascript", "text/css", - } - if content_type in _DANGEROUS_MIME_TYPES: - content_type = "application/octet-stream" + # User-controlled asset content must not render inline in the app origin + # (stored XSS via SVG/HTML/XML). Force dangerous types to download and + # override any requested inline disposition; SVG loaded into an is + # exempt, see renders_safely_as_image. Centralised through folder_paths so + # this can't drift from /view and /userdata (the previous inline set here + # omitted image/svg+xml and missed the charset/casing/+xml-dialect bypasses). + extra_headers = {} + sec_fetch_dest = request.headers.get("Sec-Fetch-Dest") + if folder_paths.is_dangerous_content_type(content_type): + # This response now depends on a request header, so it must not be + # reused across destinations by a browser or intermediary cache: an + # inline SVG primed by an fetch and replayed to a document + # navigation of the same URL would re-enable the stored XSS. + extra_headers["Vary"] = "Sec-Fetch-Dest" + extra_headers["Cache-Control"] = "no-store" + if not folder_paths.renders_safely_as_image(content_type, sec_fetch_dest): + content_type = "application/octet-stream" + disposition = "attachment" + + # mime_type is uploader-supplied and unvalidated, so it can carry + # parameters. aiohttp rejects a charset in the content_type argument with + # ValueError, which would turn a valid inline SVG into a 500. + content_type = content_type.split(";", 1)[0].strip() or "application/octet-stream" safe_name = (filename or "").replace("\r", "").replace("\n", "") encoded = urllib.parse.quote(safe_name) @@ -344,6 +370,7 @@ async def download_asset_content(request: web.Request) -> web.Response: "Content-Disposition": cd, "Content-Length": str(file_size), "X-Content-Type-Options": "nosniff", + **extra_headers, }, ) @@ -416,17 +443,6 @@ async def upload_asset(request: web.Request) -> web.Response: 400, "INVALID_BODY", f"Validation failed: {ve.json()}" ) - if spec.tags and spec.tags[0] == "models": - if ( - len(spec.tags) < 2 - or spec.tags[1] not in folder_paths.folder_names_and_paths - ): - delete_temp_file_if_exists(parsed.tmp_path) - category = spec.tags[1] if len(spec.tags) >= 2 else "" - return _build_error_response( - 400, "INVALID_BODY", f"unknown models category '{category}'" - ) - try: # Fast path: hash exists, create AssetReference without writing anything if spec.hash and parsed.provided_hash_exists is True: @@ -470,7 +486,7 @@ async def upload_asset(request: web.Request) -> web.Response: return _build_error_response(400, e.code, str(e)) except ValueError as e: delete_temp_file_if_exists(parsed.tmp_path) - return _build_error_response(400, "BAD_REQUEST", str(e)) + return _build_error_response(400, "INVALID_BODY", str(e)) except HashMismatchError as e: delete_temp_file_if_exists(parsed.tmp_path) return _build_error_response(400, "HASH_MISMATCH", str(e)) diff --git a/app/assets/api/schemas_in.py b/app/assets/api/schemas_in.py index af666746d..38a942b7b 100644 --- a/app/assets/api/schemas_in.py +++ b/app/assets/api/schemas_in.py @@ -140,7 +140,7 @@ class CreateFromHashBody(BaseModel): if v is None: return [] if isinstance(v, list): - out = [str(t).strip().lower() for t in v if str(t).strip()] + out = [str(t).strip() for t in v if str(t).strip()] seen = set() dedup = [] for t in out: @@ -149,7 +149,7 @@ class CreateFromHashBody(BaseModel): dedup.append(t) return dedup if isinstance(v, str): - return [t.strip().lower() for t in v.split(",") if t.strip()] + return list(dict.fromkeys(t.strip() for t in v.split(",") if t.strip())) return [] @@ -206,7 +206,7 @@ class TagsListQuery(BaseModel): if v is None: return v v = v.strip() - return v.lower() or None + return v or None class TagsAdd(BaseModel): @@ -220,7 +220,7 @@ class TagsAdd(BaseModel): for t in v: if not isinstance(t, str): raise TypeError("tags must be strings") - tnorm = t.strip().lower() + tnorm = t.strip() if tnorm: out.append(tnorm) seen = set() @@ -239,8 +239,8 @@ class TagsRemove(TagsAdd): class UploadAssetSpec(BaseModel): """Upload Asset operation. - - tags: optional list; if provided, first is root ('models'|'input'|'output'); - if root == 'models', second must be a valid category + - tags: labels plus one destination role ('models'|'input'|'output') for new bytes; + if role == 'models', exactly one model_type: tag is required - name: display name - user_metadata: arbitrary JSON object (optional) - hash: optional canonical 'blake3:' for validation / fast-path @@ -309,7 +309,7 @@ class UploadAssetSpec(BaseModel): norm = [] seen = set() for t in items: - tnorm = str(t).strip().lower() + tnorm = str(t).strip() if tnorm and tnorm not in seen: seen.add(tnorm) norm.append(tnorm) @@ -335,14 +335,4 @@ class UploadAssetSpec(BaseModel): @model_validator(mode="after") def _validate_order(self): - if not self.tags: - raise ValueError("at least one tag is required for uploads") - root = self.tags[0] - if root not in {"models", "input", "output"}: - raise ValueError("first tag must be one of: models, input, output") - if root == "models": - if len(self.tags) < 2: - raise ValueError( - "models uploads require a category tag as the second tag" - ) return self diff --git a/app/assets/api/schemas_out.py b/app/assets/api/schemas_out.py index 4e38e19d1..da8251499 100644 --- a/app/assets/api/schemas_out.py +++ b/app/assets/api/schemas_out.py @@ -9,8 +9,20 @@ class Asset(BaseModel): ``id`` here is the AssetReference id, not the content-addressed Asset id.""" id: str - name: str + name: str = Field( + ..., + deprecated=True, + description="Reference label, often caller-provided or derived from the filename. Deprecated for storage path/display semantics; use `loader_path` and `display_name` when present.", + ) hash: str | None = None + loader_path: str | None = Field( + default=None, + description="The value a loader consumes to load this asset. `None` when no loader can resolve the file.", + ) + display_name: str | None = Field( + default=None, + description="Human-facing label for the asset. Not unique.", + ) asset_hash: str | None = None size: int | None = None mime_type: str | None = None diff --git a/app/assets/api/upload.py b/app/assets/api/upload.py index 13d3d372c..2979f0e20 100644 --- a/app/assets/api/upload.py +++ b/app/assets/api/upload.py @@ -140,7 +140,6 @@ async def parse_multipart_upload( provided_mime_type = ((await field.text()) or "").strip() or None elif fname == "preview_id": provided_preview_id = ((await field.text()) or "").strip() or None - if not file_present and not (provided_hash and provided_hash_exists): raise UploadError( 400, "MISSING_FILE", "Form must include a 'file' part or a known 'hash'." diff --git a/app/assets/database/models.py b/app/assets/database/models.py index 9b61d309a..329cd483d 100644 --- a/app/assets/database/models.py +++ b/app/assets/database/models.py @@ -76,6 +76,8 @@ class AssetReference(Base): # Cache state fields (from former AssetCacheState) file_path: Mapped[str | None] = mapped_column(Text, nullable=True) + # In-root loader path derived from file_path at scan/ingest time. + loader_path: Mapped[str | None] = mapped_column(Text, nullable=True) mtime_ns: Mapped[int | None] = mapped_column(BigInteger, nullable=True) needs_verify: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) is_missing: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False) diff --git a/app/assets/database/queries/asset_reference.py b/app/assets/database/queries/asset_reference.py index 792411800..967b0e43a 100644 --- a/app/assets/database/queries/asset_reference.py +++ b/app/assets/database/queries/asset_reference.py @@ -650,6 +650,7 @@ def upsert_reference( name: str, mtime_ns: int, owner_id: str = "", + loader_path: str | None = None, ) -> tuple[bool, bool]: """Upsert a reference by file_path. Returns (created, updated). @@ -659,6 +660,7 @@ def upsert_reference( vals = { "asset_id": asset_id, "file_path": file_path, + "loader_path": loader_path, "name": name, "owner_id": owner_id, "mtime_ns": int(mtime_ns), @@ -686,13 +688,14 @@ def upsert_reference( AssetReference.asset_id != asset_id, AssetReference.mtime_ns.is_(None), AssetReference.mtime_ns != int(mtime_ns), + AssetReference.loader_path.is_distinct_from(loader_path), AssetReference.is_missing == True, # noqa: E712 AssetReference.deleted_at.isnot(None), ) ) .values( - asset_id=asset_id, mtime_ns=int(mtime_ns), is_missing=False, - deleted_at=None, updated_at=now, + asset_id=asset_id, mtime_ns=int(mtime_ns), loader_path=loader_path, + is_missing=False, deleted_at=None, updated_at=now, ) ) res2 = session.execute(upd) diff --git a/app/assets/database/queries/tags.py b/app/assets/database/queries/tags.py index d41d73a10..148f34801 100644 --- a/app/assets/database/queries/tags.py +++ b/app/assets/database/queries/tags.py @@ -265,6 +265,8 @@ def list_tags_with_usage( order: str = "count_desc", owner_id: str = "", ) -> tuple[list[tuple[str, str, int]], int]: + prefix_filter = prefix.strip() if prefix else "" + counts_sq = ( select( AssetReferenceTag.tag_name.label("tag_name"), @@ -293,9 +295,8 @@ def list_tags_with_usage( .join(counts_sq, counts_sq.c.tag_name == Tag.name, isouter=True) ) - if prefix: - escaped, esc = escape_sql_like_string(prefix.strip().lower()) - q = q.where(Tag.name.like(escaped + "%", escape=esc)) + if prefix_filter: + q = q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter) if not include_zero: q = q.where(func.coalesce(counts_sq.c.cnt, 0) > 0) @@ -306,9 +307,8 @@ def list_tags_with_usage( q = q.order_by(func.coalesce(counts_sq.c.cnt, 0).desc(), Tag.name.asc()) total_q = select(func.count()).select_from(Tag) - if prefix: - escaped, esc = escape_sql_like_string(prefix.strip().lower()) - total_q = total_q.where(Tag.name.like(escaped + "%", escape=esc)) + if prefix_filter: + total_q = total_q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter) if not include_zero: visible_tags_sq = ( select(AssetReferenceTag.tag_name) diff --git a/app/assets/helpers.py b/app/assets/helpers.py index 3798f3933..87734d0dc 100644 --- a/app/assets/helpers.py +++ b/app/assets/helpers.py @@ -41,10 +41,10 @@ def get_utc_now() -> datetime: def normalize_tags(tags: list[str] | None) -> list[str]: """ Normalize a list of tags by: - - Stripping whitespace and converting to lowercase. - - Removing duplicates. + - Stripping whitespace. + - Removing exact duplicates while preserving order and case. """ - return list(dict.fromkeys(t.strip().lower() for t in (tags or []) if (t or "").strip())) + return list(dict.fromkeys(t.strip() for t in (tags or []) if (t or "").strip())) def validate_blake3_hash(s: str) -> str: diff --git a/app/assets/scanner.py b/app/assets/scanner.py index 2c1e97840..42c4c1e9d 100644 --- a/app/assets/scanner.py +++ b/app/assets/scanner.py @@ -36,7 +36,7 @@ from app.assets.services.hashing import HashCheckpoint, compute_blake3_hash from app.assets.services.image_dimensions import extract_image_dimensions from app.assets.services.metadata_extract import extract_file_metadata from app.assets.services.path_utils import ( - compute_relative_filename, + compute_loader_path, get_comfy_models_folders, get_name_and_tags_from_asset_path, ) @@ -63,7 +63,7 @@ RootType = Literal["models", "input", "output"] def get_prefixes_for_root(root: RootType) -> list[str]: if root == "models": bases: list[str] = [] - for _bucket, paths in get_comfy_models_folders(): + for _bucket, paths, _exts in get_comfy_models_folders(): bases.extend(paths) return [os.path.abspath(p) for p in bases] if root == "input": @@ -81,7 +81,7 @@ def get_all_known_prefixes() -> list[str]: def collect_models_files() -> list[str]: out: list[str] = [] - for folder_name, bases in get_comfy_models_folders(): + for folder_name, bases, _exts in get_comfy_models_folders(): rel_files = folder_paths.get_filename_list(folder_name) or [] for rel_path in rel_files: if not all(is_visible(part) for part in Path(rel_path).parts): @@ -308,7 +308,7 @@ def build_asset_specs( if not stat_p.st_size: continue name, tags = get_name_and_tags_from_asset_path(abs_p) - rel_fname = compute_relative_filename(abs_p) + rel_fname = compute_loader_path(abs_p) # Extract metadata (tier 1: filesystem, tier 2: safetensors header) metadata = None @@ -430,7 +430,7 @@ def enrich_asset( return new_level initial_mtime_ns = get_mtime_ns(stat_p) - rel_fname = compute_relative_filename(file_path) + rel_fname = compute_loader_path(file_path) mime_type: str | None = None metadata = None diff --git a/app/assets/services/asset_management.py b/app/assets/services/asset_management.py index d4e4fc61c..a4c8b5a75 100644 --- a/app/assets/services/asset_management.py +++ b/app/assets/services/asset_management.py @@ -38,7 +38,7 @@ from app.assets.database.queries import ( update_reference_updated_at, ) from app.assets.helpers import select_best_live_path -from app.assets.services.path_utils import compute_relative_filename +from app.assets.services.path_utils import compute_loader_path from app.assets.services.schemas import ( AssetData, AssetDetailResult, @@ -91,7 +91,7 @@ def update_asset_metadata( update_reference_name(session, reference_id=reference_id, name=name) touched = True - computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None + computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None new_meta: dict | None = None if user_metadata is not None: diff --git a/app/assets/services/bulk_ingest.py b/app/assets/services/bulk_ingest.py index 67aad838f..c98658bf1 100644 --- a/app/assets/services/bulk_ingest.py +++ b/app/assets/services/bulk_ingest.py @@ -56,6 +56,7 @@ class ReferenceRow(TypedDict): id: str asset_id: str file_path: str + loader_path: str | None mtime_ns: int owner_id: str name: str @@ -134,6 +135,14 @@ def batch_insert_seed_assets( for spec in specs: absolute_path = os.path.abspath(spec["abs_path"]) + existing_asset_id = path_to_asset_id.get(absolute_path) + if existing_asset_id is not None: + existing_tags = asset_id_to_ref_data[existing_asset_id]["tags"] + asset_id_to_ref_data[existing_asset_id]["tags"] = list( + dict.fromkeys([*existing_tags, *spec["tags"]]) + ) + continue + asset_id = str(uuid.uuid4()) reference_id = str(uuid.uuid4()) absolute_path_list.append(absolute_path) @@ -164,6 +173,8 @@ def batch_insert_seed_assets( "id": reference_id, "asset_id": asset_id, "file_path": absolute_path, + # spec["fname"] is compute_loader_path(abs_path) from build_asset_specs. + "loader_path": spec["fname"], "mtime_ns": spec["mtime_ns"], "owner_id": owner_id, "name": spec["info_name"], diff --git a/app/assets/services/ingest.py b/app/assets/services/ingest.py index 3b6dc237c..1ffb3d634 100644 --- a/app/assets/services/ingest.py +++ b/app/assets/services/ingest.py @@ -33,8 +33,9 @@ from app.assets.services.bulk_ingest import batch_insert_seed_assets from app.assets.services.file_utils import get_size_and_mtime_ns from app.assets.services.image_dimensions import extract_image_dimensions from app.assets.services.path_utils import ( - compute_relative_filename, + compute_loader_path, get_name_and_tags_from_asset_path, + get_path_derived_tags_from_path, resolve_destination_from_tags, validate_path_within_base, ) @@ -91,6 +92,7 @@ def _ingest_file_from_path( name=info_name or os.path.basename(locator), mtime_ns=mtime_ns, owner_id=owner_id, + loader_path=compute_loader_path(locator), ) # Get the reference we just created/updated @@ -101,17 +103,32 @@ def _ingest_file_from_path( if preview_id and ref.preview_id != preview_id: ref.preview_id = preview_id - norm = normalize_tags(list(tags)) - if norm: + try: + backend_tags = get_path_derived_tags_from_path(locator) + except ValueError: + backend_tags = [] + caller_tags = normalize_tags(tags) + backend_tags = normalize_tags(backend_tags) + all_tags = normalize_tags([*caller_tags, *backend_tags]) + if all_tags: if require_existing_tags: - validate_tags_exist(session, norm) - add_tags_to_reference( - session, - reference_id=reference_id, - tags=norm, - origin=tag_origin, - create_if_missing=not require_existing_tags, - ) + validate_tags_exist(session, all_tags) + if backend_tags: + add_tags_to_reference( + session, + reference_id=reference_id, + tags=backend_tags, + origin="automatic", + create_if_missing=not require_existing_tags, + ) + if caller_tags: + add_tags_to_reference( + session, + reference_id=reference_id, + tags=caller_tags, + origin=tag_origin, + create_if_missing=not require_existing_tags, + ) _update_metadata_with_filename( session, @@ -228,7 +245,7 @@ def ingest_existing_file( "mtime_ns": mtime_ns, "info_name": name, "tags": tags, - "fname": os.path.basename(abs_path), + "fname": compute_loader_path(abs_path), "metadata": None, "hash": None, "mime_type": mime_type, @@ -288,7 +305,7 @@ def _register_existing_asset( return result new_meta = dict(user_metadata) - computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None + computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None if computed_filename: new_meta["filename"] = computed_filename @@ -335,7 +352,7 @@ def _update_metadata_with_filename( current_metadata: dict | None, user_metadata: dict[str, Any], ) -> None: - computed_filename = compute_relative_filename(file_path) if file_path else None + computed_filename = compute_loader_path(file_path) if file_path else None current_meta = current_metadata or {} new_meta = dict(current_meta) @@ -474,6 +491,10 @@ def upload_from_temp_path( existing = get_asset_by_hash(session, asset_hash=asset_hash) if existing is not None: + # Once content is already known, duplicate byte uploads are treated as + # reference-only creation. Request tags are labels only here: do not + # require upload destination tags, do not move bytes, and do not + # synthesize path-derived classification or uploaded provenance. with contextlib.suppress(Exception): if temp_path and os.path.exists(temp_path): os.remove(temp_path) @@ -535,7 +556,7 @@ def upload_from_temp_path( owner_id=owner_id, preview_id=preview_id, user_metadata=user_metadata or {}, - tags=tags, + tags=[*(tags or []), "uploaded"], tag_origin="manual", require_existing_tags=False, ) @@ -569,15 +590,19 @@ def register_file_in_place( ) -> UploadResult: """Register an already-saved file in the asset database without moving it. - Tags are derived from the filesystem path (root category + subfolder names), - merged with any caller-provided tags, matching the behavior of the scanner. + This helper is used by upload paths that have already written bytes before + registering the file, so it records the same ``uploaded`` tag as the + multipart byte-upload path. + + Tags are derived from trusted filesystem classification and merged with any + caller-provided tags, matching the behavior of the scanner. If the path is not under a known root, only the caller-provided tags are used. """ try: _, path_tags = get_name_and_tags_from_asset_path(abs_path) except ValueError: path_tags = [] - merged_tags = normalize_tags([*path_tags, *tags]) + merged_tags = normalize_tags([*path_tags, *tags, "uploaded"]) try: digest, _ = hashing.compute_blake3_hash(abs_path) diff --git a/app/assets/services/path_utils.py b/app/assets/services/path_utils.py index 892140ffb..7c27c8878 100644 --- a/app/assets/services/path_utils.py +++ b/app/assets/services/path_utils.py @@ -3,59 +3,66 @@ from pathlib import Path from typing import Literal import folder_paths -from app.assets.helpers import normalize_tags -_NON_MODEL_FOLDER_NAMES = frozenset({"custom_nodes"}) +_NON_MODEL_FOLDER_NAMES = frozenset({"configs", "custom_nodes"}) +_KNOWN_SUBFOLDER_TAGS = frozenset({"3d", "pasted", "painter", "threed", "webcam"}) -def get_comfy_models_folders() -> list[tuple[str, list[str]]]: - """Build list of (folder_name, base_paths[]) for all model locations. +def get_comfy_models_folders() -> list[tuple[str, list[str], set[str]]]: + """Build list of (folder_name, base_paths[], extensions) for all model locations. Includes every category registered in folder_names_and_paths, regardless of whether its paths are under the main models_dir, - but excludes non-model entries like custom_nodes. + but excludes non-model entries like configs and custom_nodes. + + An empty extensions set means the category accepts any extension, + matching folder_paths.filter_files_extensions semantics. """ - targets: list[tuple[str, list[str]]] = [] + targets: list[tuple[str, list[str], set[str]]] = [] for name, values in folder_paths.folder_names_and_paths.items(): if name in _NON_MODEL_FOLDER_NAMES: continue - paths, _exts = values[0], values[1] + paths, exts = values[0], values[1] if paths: - targets.append((name, paths)) + targets.append((name, paths, set(exts))) return targets def resolve_destination_from_tags(tags: list[str]) -> tuple[str, list[str]]: - """Validates and maps tags -> (base_dir, subdirs_for_fs)""" - if not tags: - raise ValueError("tags must not be empty") - root = tags[0].lower() + """Validates and maps upload routing tags -> (base_dir, subdirs_for_fs). + + The request tags are only used to choose the write destination. Extra tags + remain labels; they do not become path components or trusted classification. + """ + destination_roles = [t for t in tags if t in {"input", "models", "output"}] + if len(destination_roles) != 1: + raise ValueError("uploads require exactly one destination role: input, models, or output") + + root = destination_roles[0] if root == "models": - if len(tags) < 2: - raise ValueError("at least two tags required for model asset") + model_type_tags = [t for t in tags if t.startswith("model_type:")] + if len(model_type_tags) != 1: + raise ValueError("models uploads require exactly one model_type: tag") + folder_name = model_type_tags[0].split(":", 1)[1] + if not folder_name: + raise ValueError("models uploads require exactly one model_type: tag") + model_folder_paths = { + name: paths for name, paths, _exts in get_comfy_models_folders() + } try: - bases = folder_paths.folder_names_and_paths[tags[1]][0] + bases = model_folder_paths[folder_name] except KeyError: - raise ValueError(f"unknown model category '{tags[1]}'") + raise ValueError(f"unknown model category '{folder_name}'") if not bases: - raise ValueError(f"no base path configured for category '{tags[1]}'") + raise ValueError(f"no base path configured for category '{folder_name}'") base_dir = os.path.abspath(bases[0]) - raw_subdirs = tags[2:] elif root == "input": base_dir = os.path.abspath(folder_paths.get_input_directory()) - raw_subdirs = tags[1:] - elif root == "output": - base_dir = os.path.abspath(folder_paths.get_output_directory()) - raw_subdirs = tags[1:] else: - raise ValueError(f"unknown root tag '{tags[0]}'; expected 'models', 'input', or 'output'") - _sep_chars = frozenset(("/", "\\", os.sep)) - for i in raw_subdirs: - if i in (".", "..") or _sep_chars & set(i): - raise ValueError("invalid path component in tags") + base_dir = os.path.abspath(folder_paths.get_output_directory()) - return base_dir, raw_subdirs if raw_subdirs else [] + return base_dir, [] def validate_path_within_base(candidate: str, base: str) -> None: @@ -65,14 +72,79 @@ def validate_path_within_base(candidate: str, base: str) -> None: raise ValueError("destination escapes base directory") -def compute_relative_filename(file_path: str) -> str | None: +def _compute_relative_path(child: str, parent: str) -> str: + rel = os.path.relpath(os.path.abspath(child), os.path.abspath(parent)) + if rel == ".": + return "" + return rel.replace(os.sep, "/") + + +def _is_relative_to(child: str, parent: str) -> bool: + return Path(os.path.abspath(child)).is_relative_to(os.path.abspath(parent)) + + +def compute_asset_response_paths(file_path: str) -> tuple[str, str | None] | None: + """Return (logical_path, display_name) for a file path. + + ``logical_path`` is the internal namespaced storage locator (e.g. + ``models/checkpoints/foo/bar.safetensors``); ``display_name`` is the + human-facing label below that namespace, served on Asset responses. These + are storage locators, not model-loader namespaces. Registered model-folder + membership is represented by backend tags such as + ``model_type:``; these paths only use known storage roots. """ - Return the model's path relative to the last well-known folder (the model category), - using forward slashes, eg: + fp_abs = os.path.abspath(file_path) + candidates: list[tuple[int, int, str, str]] = [] + + for order, (namespace, base) in enumerate( + ( + ("input", folder_paths.get_input_directory()), + ("output", folder_paths.get_output_directory()), + ("temp", folder_paths.get_temp_directory()), + ("models", getattr(folder_paths, "models_dir", "")), + ) + ): + if not base: + continue + base_abs = os.path.abspath(base) + if _is_relative_to(fp_abs, base_abs): + candidates.append((len(base_abs), -order, namespace, base_abs)) + + if not candidates: + return None + + _base_len, _order, namespace, base = max(candidates) + rel = _compute_relative_path(fp_abs, base) + public_path = f"{namespace}/{rel}" if rel else namespace + return public_path, rel or None + + +def compute_display_name(file_path: str) -> str | None: + """Return the asset's `display_name`, or None for unknown paths.""" + result = compute_asset_response_paths(file_path) + return result[1] if result else None + + +def compute_logical_path(file_path: str) -> str | None: + """Return the internal namespaced storage locator, or None for unknown paths.""" + result = compute_asset_response_paths(file_path) + return result[0] if result else None + + +def compute_loader_path(file_path: str) -> str | None: + """ + Return the asset's in-root loader path: the path relative to the last + well-known folder (the model category), using forward slashes, eg: /.../models/checkpoints/flux/123/flux.safetensors -> "flux/123/flux.safetensors" /.../models/text_encoders/clip_g.safetensors -> "clip_g.safetensors" - For non-model paths, returns None. + This is the value model loaders consume (the model category is dropped). It + is persisted as ``AssetReference.loader_path`` and served as the public + Asset response `loader_path` field. The human-facing `display_name` comes + from compute_asset_response_paths(). + + For input/output/temp paths the full path relative to that root is returned. + For paths outside any known root, returns None. """ try: root_category, rel_path = get_asset_category_and_relative_path(file_path) @@ -116,9 +188,10 @@ def get_asset_category_and_relative_path( def _compute_relative(child: str, parent: str) -> str: # Normalize relative path, stripping any leading ".." components # by anchoring to root (os.sep) then computing relpath back from it. - return os.path.relpath( + rel = os.path.relpath( os.path.join(os.sep, os.path.relpath(child, parent)), os.sep ) + return "" if rel == "." else rel.replace(os.sep, "/") # 1) input input_base = os.path.abspath(folder_paths.get_input_directory()) @@ -136,8 +209,14 @@ def get_asset_category_and_relative_path( return "temp", _compute_relative(fp_abs, temp_base) # 4) models (check deepest matching base to avoid ambiguity) + ext = os.path.splitext(fp_abs)[1].lower() best: tuple[int, str, str] | None = None # (base_len, bucket, rel_inside_bucket) - for bucket, bases in get_comfy_models_folders(): + for bucket, bases, extensions in get_comfy_models_folders(): + # A bucket only lists files within its extension set (empty set + # accepts any extension), so a bucket that cannot load the file + # must not contribute a loader path. + if extensions and ext not in extensions: + continue for b in bases: base_abs = os.path.abspath(b) if not _check_is_within(fp_abs, base_abs): @@ -149,25 +228,111 @@ def get_asset_category_and_relative_path( if best is not None: _, bucket, rel_inside = best combined = os.path.join(bucket, rel_inside) - return "models", os.path.relpath(os.path.join(os.sep, combined), os.sep) + normalized = os.path.relpath(os.path.join(os.sep, combined), os.sep) + return "models", normalized.replace(os.sep, "/") raise ValueError( f"Path is not within input, output, temp, or configured model bases: {file_path}" ) +def get_backend_system_tags_from_path(path: str) -> list[str]: + """Return trusted backend tags derived from current filesystem facts. + + The returned tags are only the backend-generated system tags: ``models``, + ``model_type:``, ``input``, ``output``, and ``temp``. Model + type tags are based on registered folder names, not path components. + + A ``model_type:`` tag is only emitted when the file's + extension is accepted by that folder's registered extension set, so + categories sharing a base directory tag only the files they can + actually load. Files under a model base whose extension matches no + category still get the ``models`` tag. + """ + fp_abs = os.path.abspath(path) + fp_path = Path(fp_abs) + tags: list[str] = [] + + def _add(tag: str) -> None: + if tag not in tags: + tags.append(tag) + + for role, base in ( + ("input", folder_paths.get_input_directory()), + ("output", folder_paths.get_output_directory()), + ("temp", folder_paths.get_temp_directory()), + ): + if fp_path.is_relative_to(os.path.abspath(base)): + _add(role) + + ext = os.path.splitext(fp_abs)[1].lower() + model_types: list[str] = [] + under_models_base = False + for folder_name, bases, extensions in get_comfy_models_folders(): + for base in bases: + if fp_path.is_relative_to(os.path.abspath(base)): + under_models_base = True + # Empty set accepts any extension, matching + # folder_paths.filter_files_extensions semantics. + if not extensions or ext in extensions: + model_types.append(folder_name) + break + + if under_models_base: + _add("models") + for folder_name in model_types: + _add(f"model_type:{folder_name}") + + if not tags: + raise ValueError( + f"Path is not within input, output, temp, or configured model bases: {path}" + ) + return tags + + +def get_known_subfolder_tags(subfolder: str | None) -> list[str]: + """Return tags for known UI/input subfolder names.""" + if subfolder in _KNOWN_SUBFOLDER_TAGS: + return [subfolder] + return [] + + +def get_known_input_subfolder_tags_from_path(path: str) -> list[str]: + """Return known input-layout tags for files in canonical input subfolders. + + These are compatibility tags for current UI-origin input directories such as + ``pasted`` and ``webcam``. They are intentionally narrow: only files directly + inside a known top-level input directory receive the matching tag. + """ + fp_abs = os.path.abspath(path) + input_base = os.path.abspath(folder_paths.get_input_directory()) + if not Path(fp_abs).is_relative_to(input_base): + return [] + + rel = os.path.relpath(fp_abs, input_base) + parts = Path(rel).parts + if len(parts) == 2: + return get_known_subfolder_tags(parts[0]) + return [] + + +def get_path_derived_tags_from_path(path: str) -> list[str]: + """Return all backend-derived tags for an asset path.""" + tags = get_backend_system_tags_from_path(path) + for tag in get_known_input_subfolder_tags_from_path(path): + if tag not in tags: + tags.append(tag) + return tags + + def get_name_and_tags_from_asset_path(file_path: str) -> tuple[str, list[str]]: """Return (name, tags) derived from a filesystem path. - name: base filename with extension - - tags: [root_category] + parent folder names in order + - tags: backend-derived tags from root/model classification and known input + subfolder layout conventions Raises: ValueError: path does not belong to any known root. """ - root_category, some_path = get_asset_category_and_relative_path(file_path) - p = Path(some_path) - parent_parts = [ - part for part in p.parent.parts if part not in (".", "..", p.anchor) - ] - return p.name, list(dict.fromkeys(normalize_tags([root_category, *parent_parts]))) + return Path(file_path).name, get_path_derived_tags_from_path(file_path) diff --git a/app/assets/services/schemas.py b/app/assets/services/schemas.py index 4d2af8a02..0fda6871d 100644 --- a/app/assets/services/schemas.py +++ b/app/assets/services/schemas.py @@ -25,6 +25,7 @@ class ReferenceData: preview_id: str | None created_at: datetime updated_at: datetime + loader_path: str | None = None system_metadata: dict[str, Any] | None = None job_id: str | None = None last_access_time: datetime | None = None @@ -93,6 +94,7 @@ def extract_reference_data(ref: AssetReference) -> ReferenceData: id=ref.id, name=ref.name, file_path=ref.file_path, + loader_path=ref.loader_path, user_metadata=ref.user_metadata, preview_id=ref.preview_id, system_metadata=ref.system_metadata, diff --git a/app/logger.py b/app/logger.py index bde815822..fe82c40c9 100644 --- a/app/logger.py +++ b/app/logger.py @@ -2,9 +2,12 @@ from collections import deque from datetime import datetime import io import logging +import os import sys import threading +import comfy.internal_logging + ANSI_NAMED_COLORS = { 'black': '\033[30m', 'red': '\033[31m', @@ -18,6 +21,7 @@ ANSI_NAMED_COLORS = { ANSI_LEVEL_COLORS = { 'DEBUG': ANSI_NAMED_COLORS['cyan'], + 'DETAIL': ANSI_NAMED_COLORS['blue'], 'INFO': ANSI_NAMED_COLORS['green'], 'WARNING': ANSI_NAMED_COLORS['yellow'], 'ERROR': ANSI_NAMED_COLORS['red'], @@ -85,7 +89,12 @@ def on_flush(callback): if stderr_interceptor is not None: stderr_interceptor.on_flush(callback) -def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool = False): + +def get_log_level(level): + return comfy.internal_logging.DETAIL if level == "DETAIL" else logging.getLevelName(level) + + +def setup_logger(log_level: str = 'INFO', file_outputs=None, capacity: int = 300, use_stdout: bool = False): global logs if logs: return @@ -99,13 +108,18 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr) # Setup default global logger + if file_outputs is None: + file_outputs = [('DETAIL', 'comfyui_detail.log')] logger = logging.getLogger() - logger.setLevel(log_level) + console_level = get_log_level(log_level) + file_levels = [get_log_level(level) for level, _ in file_outputs] + logger.setLevel(min([console_level, *file_levels])) formatter = ColoredFormatter("%(message)s") stream_handler = logging.StreamHandler() stream_handler.setFormatter(formatter) + stream_handler.setLevel(console_level) if use_stdout: # Only errors and critical to stderr @@ -114,11 +128,24 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300, use_stdout: bool # Lesser to stdout stdout_handler = logging.StreamHandler(sys.stdout) stdout_handler.setFormatter(formatter) + stdout_handler.setLevel(console_level) stdout_handler.addFilter(lambda record: record.levelno < logging.ERROR) logger.addHandler(stdout_handler) logger.addHandler(stream_handler) + for output_level, output_path in file_outputs: + output_path = os.path.abspath(output_path) + try: + output_handler = logging.FileHandler(output_path, encoding="utf-8") + except OSError as e: + logging.warning("Could not open %s log %s: %s", output_level, output_path, e) + continue + output_handler.setLevel(get_log_level(output_level)) + output_handler.setFormatter(logging.Formatter("[%(asctime)s] [%(levelname)s] %(message)s")) + logger.addHandler(output_handler) + logging.info("%s log: %s", output_level.title(), output_path) + STARTUP_WARNINGS = [] diff --git a/app/model_manager.py b/app/model_manager.py index 8f6e34b33..5928781ca 100644 --- a/app/model_manager.py +++ b/app/model_manager.py @@ -35,7 +35,11 @@ class ModelFileManager: for folder in model_types: if folder in folder_black_list: continue - output_folders.append({"name": folder, "folders": folder_paths.get_folder_paths(folder)}) + output_folders.append({ + "name": folder, + "folders": folder_paths.get_folder_paths(folder), + "extensions": sorted(folder_paths.folder_names_and_paths[folder][1]), + }) return web.json_response(output_folders) # NOTE: This is an experiment to replace `/models/{folder}` @@ -50,21 +54,45 @@ class ModelFileManager: @routes.get("/experiment/models/preview/{folder}/{path_index}/{filename:.*}") async def get_model_preview(request): folder_name = request.match_info.get("folder", None) - path_index = int(request.match_info.get("path_index", None)) filename = request.match_info.get("filename", None) if folder_name not in folder_paths.folder_names_and_paths: return web.Response(status=404) + # The "{filename:.*}" capture also matches the empty string, which + # would resolve to the folder itself; reject it explicitly. + if not filename: + return web.Response(status=400) + + try: + path_index = int(request.match_info.get("path_index", None)) + except (TypeError, ValueError): + return web.Response(status=400) + folders = folder_paths.folder_names_and_paths[folder_name] + if path_index < 0 or path_index >= len(folders[0]): + return web.Response(status=404) folder = folders[0][path_index] - full_filename = os.path.join(folder, filename) + full_filename = os.path.normpath(os.path.join(folder, filename)) + + # Prevent path traversal: the requested file must stay within the + # configured model folder. `filename` is an unrestricted ".*" capture, + # so values like "../../../../etc/passwd" would otherwise escape it. + if not folder_paths.is_within_directory(folder, full_filename): + return web.Response(status=403) previews = self.get_model_previews(full_filename) default_preview = previews[0] if len(previews) > 0 else None if default_preview is None or (isinstance(default_preview, str) and not os.path.isfile(default_preview)): return web.Response(status=404) + # The preview is selected by a glob inside get_model_previews, so a + # companion file (e.g. "model.preview.png") could itself be a symlink + # resolving outside the model folder. Re-validate the file actually + # opened: is_within_directory realpaths it, catching symlink escape. + if isinstance(default_preview, str) and not folder_paths.is_within_directory(folder, default_preview): + return web.Response(status=403) + try: with Image.open(default_preview) as img: img_bytes = BytesIO() diff --git a/app/user_manager.py b/app/user_manager.py index 7b11e381c..55e7e81e3 100644 --- a/app/user_manager.py +++ b/app/user_manager.py @@ -6,6 +6,7 @@ import glob import shutil import logging import tempfile +import mimetypes from aiohttp import web from urllib import parse from comfy.cli_args import args @@ -336,7 +337,29 @@ class UserManager(): if not isinstance(path, str): return path - return web.FileResponse(path) + # User data files are arbitrary user-supplied content and are never + # meant to render inline. Disable MIME sniffing and force a download + # so uploaded markup/scripts can't execute in the app origin (stored + # XSS). Content-Disposition: attachment is the load-bearing guard; + # the content-type override and nosniff are defence in depth. + content_type = mimetypes.guess_type(path)[0] or 'application/octet-stream' + + user_root = self.get_request_user_filepath(request, None, create_dir=False) + is_user_css = path == os.path.abspath(os.path.join(user_root, "user.css")) + + if is_user_css: + content_type = "text/css" + disposition = "inline" + else: + if folder_paths.is_dangerous_content_type(content_type): + content_type = 'application/octet-stream' + disposition = "attachment" + + return web.FileResponse(path, headers={ + "Content-Type": content_type, + "X-Content-Type-Options": "nosniff", + "Content-Disposition": disposition, + }) @routes.post("/userdata/{file}") async def post_userdata(request): diff --git a/comfy/cli_args.py b/comfy/cli_args.py index e3099a230..ee9e1ce9f 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -33,6 +33,31 @@ class EnumAction(argparse.Action): setattr(namespace, self.dest, value) +LOG_LEVELS = ('DEBUG', 'DETAIL', 'INFO', 'WARNING', 'ERROR', 'CRITICAL') + + +class VerboseAction(argparse.Action): + def __call__(self, parser, namespace, values, option_string=None): + if len(values) == 0: + output = ('DEBUG', None) + elif len(values) == 1 and values[0] in LOG_LEVELS: + output = (values[0], None) + elif len(values) == 2 and values[0] in LOG_LEVELS: + output = tuple(values) + else: + parser.error(f"{option_string} expects no values, a console LEVEL, or LEVEL FILE") + setattr(namespace, self.dest, [*getattr(namespace, self.dest, []), output]) + + +def get_console_log_level(outputs): + console_levels = [level for level, path in outputs if path is None] + return min(console_levels, key=LOG_LEVELS.index, default='INFO') + + +def get_file_log_outputs(outputs): + return [(level, path) for level, path in outputs if path is not None] + + parser = argparse.ArgumentParser() parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)") @@ -92,6 +117,7 @@ parser.add_argument("--directml", type=int, nargs="?", metavar="DIRECTML_DEVICE" parser.add_argument("--oneapi-device-selector", type=str, default=None, metavar="SELECTOR_STRING", help="Sets the oneAPI device(s) this instance will use.") parser.add_argument("--supports-fp8-compute", action="store_true", help="ComfyUI will act like if the device supports fp8 compute.") parser.add_argument("--enable-triton-backend", action="store_true", help="ComfyUI will enable the use of Triton backend in comfy-kitchen. Is disabled at launch by default.") +parser.add_argument("--disable-triton-backend", action="store_true", help="Force-disable the comfy-kitchen Triton backend, overriding the automatic ROCm/AMD default and --enable-triton-backend.") class LatentPreviewMethod(enum.Enum): NoPreviews = "none" @@ -111,7 +137,7 @@ parser.add_argument("--preview-method", type=LatentPreviewMethod, default=Latent parser.add_argument("--preview-size", type=int, default=512, help="Sets the maximum preview size for sampler nodes.") cache_group = parser.add_mutually_exclusive_group() -cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 96GB).") +cache_group.add_argument("--cache-ram", nargs='*', type=float, default=[], metavar="GB", help="Use RAM pressure caching with the specified headroom thresholds. This is the default caching mode. The first value sets the active-cache threshold; the optional second value sets the inactive-cache/pin threshold. Defaults when no values are provided: active 10%% of system RAM (min 2GB, max 10GB), inactive 100%% of system RAM (max 128GB).") cache_group.add_argument("--cache-classic", action="store_true", help="Use the old style (aggressive) caching.") cache_group.add_argument("--cache-lru", type=int, default=0, help="Use LRU caching with a maximum of N node results cached. May use more RAM/VRAM.") cache_group.add_argument("--cache-none", action="store_true", help="Reduced RAM/VRAM usage at the expense of executing every node for each run.") @@ -146,6 +172,7 @@ vram_group.add_argument("--cpu", action="store_true", help="To use the CPU for e parser.add_argument("--reserve-vram", type=float, default=None, help="Set the amount of vram in GB you want to reserve for use by your OS/other software. By default some amount is reserved depending on your OS.") parser.add_argument("--vram-headroom", type=float, default=0, help="Set the amount of vram in GB for DynamicVRAM to maintain as extra headroom above default. ComfyUI will try and keep this much VRAM completely free and unused, even counting VRAM from other apps.") +parser.add_argument("--disable-nvml-pressure", action="store_true", help="Use CUDA instead of NVML for DynamicVRAM memory pressure.") parser.add_argument("--async-offload", nargs='?', const=2, type=int, default=None, metavar="NUM_STREAMS", help="Use async weight offloading. An optional argument controls the amount of offload streams. Default is 2. Enabled by default on Nvidia.") parser.add_argument("--disable-async-offload", action="store_true", help="Disable async weight offloading.") @@ -186,7 +213,7 @@ parser.add_argument("--disable-api-nodes", action="store_true", help="Disable lo parser.add_argument("--multi-user", action="store_true", help="Enables per-user storage.") -parser.add_argument("--verbose", default='INFO', const='DEBUG', nargs="?", choices=['DEBUG', 'INFO', 'WARNING', 'ERROR', 'CRITICAL'], help='Set the logging level') +parser.add_argument("--verbose", action=VerboseAction, nargs='*', default=[], metavar='LEVEL FILE', help='Set console logging with no values or LEVEL, or add a LEVEL FILE log output. May be repeated.') parser.add_argument("--log-stdout", action="store_true", help="Send normal process output to stdout instead of stderr (default).") @@ -225,6 +252,7 @@ parser.add_argument( ) parser.add_argument("--user-directory", type=is_valid_directory, default=None, help="Set the ComfyUI user directory with an absolute path. Overrides --base-directory.") +parser.add_argument("--models-directory", type=is_valid_directory, default=None, help="Set the ComfyUI models directory. Overrides the models folder in --base-directory.") parser.add_argument("--enable-compress-response-body", action="store_true", help="Enable compressing response body.") @@ -240,6 +268,7 @@ database_default_path = os.path.abspath( ) parser.add_argument("--database-url", type=str, default=f"sqlite:///{database_default_path}", help="Specify the database URL, e.g. for an in-memory database you can use 'sqlite:///:memory:'.") parser.add_argument("--enable-assets", action="store_true", help="Enable the assets system (API routes, database synchronization, and background scanning).") +parser.add_argument("--enable-asset-hashing", action="store_true", help="Compute blake3 content hashes when scanning assets. Hashing enables future asset-portability features (deduplication, cross-machine model resolution) but adds startup cost and per-output cost on large models directories. Off by default; enable to opt in.") parser.add_argument("--feature-flag", type=str, action='append', default=[], metavar="KEY[=VALUE]", help="Set a server feature flag. Use KEY=VALUE to set an explicit value, or bare KEY to set it to true. Can be specified multiple times. Boolean values (true/false) and numbers are auto-converted. Examples: --feature-flag show_signin_button=true or --feature-flag show_signin_button") parser.add_argument("--list-feature-flags", action="store_true", help="Print the registry of known CLI-settable feature flags as JSON and exit.") diff --git a/comfy/comfy_api_env.py b/comfy/comfy_api_env.py new file mode 100644 index 000000000..17b47933f --- /dev/null +++ b/comfy/comfy_api_env.py @@ -0,0 +1,46 @@ +"""Runtime config the frontend reads from /features to follow --comfy-api-base. + +For a non-prod comfy.org backend (staging or an ephemeral preview env), "/features" exposes the api and +platform base so the frontend talks to it without a rebuild, plus the Firebase environment it should use. +Prod bases are left alone and keep their build-time defaults. +""" + +from typing import Any +from urllib.parse import urlparse + +from comfy.cli_args import args + +_STAGING_API_HOST = "stagingapi.comfy.org" +_TESTENV_HOST_SUFFIX = ".testenvs.comfy.org" +_STAGING_PLATFORM_BASE_URL = "https://stagingplatform.comfy.org" + + +def _is_staging_tier(host: str) -> bool: + return host == _STAGING_API_HOST or host.endswith(_TESTENV_HOST_SUFFIX) + + +def normalize_comfy_api_base(url: str) -> str: + """Rewrite a testenv's friendly main host to its comfy-api '-registry' sibling.""" + parsed = urlparse(url) + host = parsed.hostname or "" + if not host.endswith(_TESTENV_HOST_SUFFIX): + return url + label = host[: -len(_TESTENV_HOST_SUFFIX)] + if label.endswith("-registry"): + return url + return f"{parsed.scheme or 'https'}://{label}-registry{_TESTENV_HOST_SUFFIX}" + + +def environment_overrides_for_base(base_url: str) -> dict[str, Any] | None: + """The /features overrides for a staging-tier base, or None for prod.""" + if not _is_staging_tier(urlparse(base_url).hostname or ""): + return None + return { + "comfy_api_base_url": normalize_comfy_api_base(base_url).rstrip("/"), + "comfy_platform_base_url": _STAGING_PLATFORM_BASE_URL, + "firebase_env": "dev", + } + + +def get_environment_overrides() -> dict[str, Any] | None: + return environment_overrides_for_base(getattr(args, "comfy_api_base", "") or "") diff --git a/comfy/internal_logging.py b/comfy/internal_logging.py new file mode 100644 index 000000000..cc785296d --- /dev/null +++ b/comfy/internal_logging.py @@ -0,0 +1,10 @@ +import logging + + +DETAIL = 15 +logging.addLevelName(DETAIL, "DETAIL") + + +def detail(message, *args, **kwargs): + kwargs.setdefault("stacklevel", 2) + logging.log(DETAIL, message, *args, **kwargs) diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index bbdfd4bc2..c4270022b 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -434,8 +434,177 @@ class LTXV(LatentFormat): class LTXAV(LTXV): def __init__(self): - self.latent_rgb_factors = None - self.latent_rgb_factors_bias = None + # video-stream preview factors for the packed AV latent (audio stream is not previewed) + self.latent_rgb_factors = [ + [ 0.001135, -0.010555, -0.004925], + [-0.008019, -0.006231, -0.005564], + [ 0.012637, 0.005605, 0.012713], + [ 0.023454, 0.020771, 0.017844], + [-0.011940, -0.000932, 0.009292], + [ 0.018602, 0.011018, 0.013969], + [-0.036369, -0.046631, -0.057898], + [-0.031919, 0.000131, 0.015214], + [ 0.014519, 0.021041, 0.015325], + [ 0.018889, 0.016149, -0.002836], + [-0.003784, -0.006057, -0.008195], + [ 0.013262, 0.030259, 0.029775], + [ 0.050465, 0.050366, 0.025255], + [ 0.018628, 0.007691, 0.002893], + [-0.015698, -0.008451, -0.000676], + [-0.013600, -0.012587, -0.004437], + [ 0.012482, 0.021469, 0.027913], + [-0.018241, -0.013488, -0.010975], + [ 0.013828, 0.012568, 0.021984], + [ 0.017911, 0.006552, 0.005567], + [ 0.026769, 0.006803, -0.009360], + [-0.006794, -0.008447, -0.013921], + [ 0.029708, 0.018671, 0.022811], + [-0.014732, -0.019169, 0.000903], + [ 0.019607, 0.032595, 0.053409], + [-0.003721, 0.003976, 0.010364], + [-0.020193, -0.026076, -0.036068], + [-0.002328, 0.006527, 0.013052], + [ 0.017171, 0.009224, 0.006548], + [ 0.001104, -0.000591, 0.000147], + [-0.000217, 0.011834, 0.017945], + [-0.015329, -0.012463, -0.006178], + [-0.009478, -0.008680, -0.004107], + [-0.005565, -0.006006, -0.001493], + [ 0.009451, 0.008794, 0.013207], + [-0.009989, -0.008027, -0.009568], + [-0.001505, -0.008805, -0.006828], + [ 0.001105, 0.008999, 0.009079], + [ 0.025935, 0.016426, 0.008036], + [ 0.006313, 0.000694, -0.006039], + [-0.001893, -0.006951, -0.009560], + [-0.007082, -0.002566, -0.007152], + [-0.005231, 0.004829, 0.008220], + [-0.004333, 0.001251, -0.004852], + [-0.017024, -0.012730, -0.007457], + [ 0.024988, 0.032963, 0.036556], + [ 0.013697, 0.012278, 0.009979], + [-0.013751, -0.008369, -0.015446], + [-0.009348, -0.001047, 0.007622], + [-0.003135, -0.003350, -0.003766], + [ 0.007436, 0.004957, 0.010480], + [ 0.018315, 0.022066, 0.021104], + [-0.005621, -0.006770, -0.008219], + [-0.007427, 0.001911, -0.001231], + [-0.007413, 0.000486, -0.006039], + [-0.014698, -0.007160, 0.006509], + [ 0.013775, 0.014185, 0.008203], + [ 0.060246, 0.069787, 0.072833], + [ 0.009861, 0.004870, 0.001194], + [-0.003660, 0.003251, 0.008015], + [ 0.003696, -0.003680, -0.008851], + [ 0.014924, 0.006196, 0.005282], + [-0.006740, -0.004319, -0.006729], + [ 0.020635, 0.015163, 0.012385], + [-0.032623, -0.006105, 0.010436], + [-0.058988, -0.030162, -0.037961], + [-0.035614, -0.021929, -0.011062], + [-0.023412, -0.011305, -0.005054], + [-0.002716, -0.005184, -0.004084], + [ 0.014591, 0.015294, 0.014045], + [ 0.008310, 0.002466, -0.003225], + [ 0.005176, 0.001119, 0.000695], + [-0.021569, -0.030886, -0.044732], + [ 0.007517, 0.003891, 0.000551], + [-0.006793, 0.004059, 0.010184], + [-0.086481, -0.082033, -0.083414], + [ 0.004192, 0.000762, -0.008658], + [ 0.010970, 0.009002, 0.007384], + [ 0.004042, -0.006732, -0.011031], + [ 0.012164, 0.006401, 0.007483], + [ 0.029252, 0.013990, 0.011128], + [ 0.048452, 0.034648, 0.016269], + [ 0.024104, 0.012647, 0.011754], + [-0.013216, -0.020192, -0.019752], + [-0.010799, -0.008535, -0.005467], + [ 0.005823, 0.001403, 0.001890], + [ 0.052393, 0.044771, 0.032777], + [ 0.007576, -0.008080, -0.012453], + [ 0.009830, 0.004244, 0.001213], + [-0.025867, -0.013169, -0.010636], + [ 0.008494, 0.003135, 0.000790], + [ 0.003969, -0.002625, -0.010204], + [ 0.006509, 0.008272, 0.020819], + [-0.004943, -0.013424, -0.015351], + [ 0.005541, 0.009136, -0.003666], + [-0.014300, -0.015864, -0.016853], + [ 0.002650, 0.028393, 0.014125], + [-0.027661, -0.045422, -0.064995], + [ 0.009220, 0.015522, 0.010574], + [-0.002236, 0.002915, 0.004557], + [-0.020269, -0.008212, -0.000532], + [ 0.019294, 0.003655, -0.002809], + [ 0.007116, -0.002784, 0.000017], + [ 0.057277, 0.073270, 0.074401], + [-0.002616, -0.001696, -0.000498], + [ 0.007248, 0.009793, 0.022829], + [-0.002590, -0.005601, -0.000436], + [-0.007681, 0.003893, -0.004119], + [-0.057392, -0.045545, -0.025290], + [ 0.045188, 0.047985, 0.054059], + [ 0.000937, -0.008861, -0.038406], + [-0.010192, -0.008036, -0.005385], + [-0.030222, -0.027498, -0.030765], + [-0.008359, 0.013247, 0.010918], + [ 0.004102, 0.002093, 0.006934], + [ 0.039461, 0.027339, 0.008284], + [-0.075747, -0.076340, -0.071625], + [ 0.002692, 0.005096, -0.002247], + [-0.002453, -0.002785, -0.010483], + [ 0.012265, 0.005481, 0.001729], + [ 0.017755, 0.008655, 0.003532], + [ 0.055560, 0.049128, 0.044137], + [-0.025861, -0.023798, -0.018815], + [-0.014876, -0.010770, -0.010713], + [-0.017315, -0.012599, -0.008661], + [-0.008461, -0.006210, -0.007744], + [-0.040175, -0.042255, -0.048119], + [-0.019355, -0.021055, -0.021919], + ] + self.latent_rgb_factors_bias = [-0.347892, -0.363814, -0.370287] + +class MiniMaxH3Video(LatentFormat): + latent_channels = 24 + latent_dimensions = 3 + spacial_downscale_ratio = 16 + temporal_downscale_ratio = 4 + scale_factor = 1.0 + + latent_rgb_factors = [ + [-0.018555, 0.024344, -0.017536], + [ 0.150164, 0.137244, 0.129221], + [ 0.027367, -0.050369, -0.208606], + [-0.000793, -0.164622, -0.323161], + [-0.048556, 0.013970, -0.074286], + [ 0.011740, 0.014172, -0.006906], + [ 0.061517, 0.061212, 0.110025], + [ 0.035321, 0.086879, 0.110059], + [-0.017426, 0.002997, 0.035356], + [ 0.531539, 0.548819, 0.624404], + [-0.024968, -0.040234, -0.034302], + [-0.032549, -0.029096, -0.017221], + [ 0.022609, 0.020286, 0.050661], + [-0.084001, -0.038131, -0.020805], + [-0.018830, 0.010412, 0.061120], + [ 0.020777, 0.011196, -0.030994], + [-0.008390, -0.012201, -0.025687], + [-0.013281, -0.002924, 0.006331], + [ 0.000260, 0.001833, -0.011038], + [ 0.105471, 0.100482, 0.132106], + [ 0.016529, 0.015213, 0.009999], + [-0.014015, -0.017438, -0.019134], + [-0.033787, -0.009984, -0.019725], + [ 0.004224, 0.017284, 0.027196], + ] + latent_rgb_factors_bias = [ 0.057426, -0.022078, -0.071449] + +class MiniMaxH3AV(MiniMaxH3Video): + # max channels across the two streams (video 24, audio 32) so per-stream slices keep both streams whole + latent_channels = 32 class HunyuanVideo(LatentFormat): latent_channels = 16 @@ -779,6 +948,10 @@ class ACEAudio(LatentFormat): latent_channels = 8 latent_dimensions = 2 +class SeedVR2(LatentFormat): + latent_channels = 16 + latent_dimensions = 3 + class ACEAudio15(LatentFormat): latent_channels = 64 latent_dimensions = 1 diff --git a/comfy/ldm/ace/ace_step15.py b/comfy/ldm/ace/ace_step15.py index 2ca2d26c4..02182c49f 100644 --- a/comfy/ldm/ace/ace_step15.py +++ b/comfy/ldm/ace/ace_step15.py @@ -217,10 +217,7 @@ class AceStepAttention(nn.Module): cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) - n_rep = self.num_heads // self.num_kv_heads - if n_rep > 1: - key_states = key_states.repeat_interleave(n_rep, dim=1) - value_states = value_states.repeat_interleave(n_rep, dim=1) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} attn_bias = None if self.sliding_window is not None and not self.is_cross_attention: @@ -244,7 +241,7 @@ class AceStepAttention(nn.Module): else: attn_bias = window_bias - attn_output = optimized_attention(query_states, key_states, value_states, self.num_heads, attn_bias, skip_reshape=True, low_precision_attention=False) + attn_output = optimized_attention(query_states, key_states, value_states, self.num_heads, attn_bias, skip_reshape=True, low_precision_attention=False, **gqa_kwargs) attn_output = self.o_proj(attn_output) return attn_output diff --git a/comfy/ldm/anima/lllite.py b/comfy/ldm/anima/lllite.py new file mode 100644 index 000000000..5c950ec89 --- /dev/null +++ b/comfy/ldm/anima/lllite.py @@ -0,0 +1,278 @@ +import re + +import torch +from torch import nn +import torch.nn.functional as F + +import comfy.ops +import comfy.utils + + +MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$") + + +def _group_norm(channels, device=None, dtype=None, operations=None): + groups = 8 + while groups > 1 and channels % groups != 0: + groups //= 2 + return operations.GroupNorm(groups, channels, device=device, dtype=dtype) + + +class AnimaLLLiteResBlock(nn.Module): + def __init__(self, channels, device=None, dtype=None, operations=None): + super().__init__() + self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations) + self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype) + self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations) + self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype) + + def forward(self, x): + h = self.conv1(F.silu(self.norm1(x))) + h = self.conv2(F.silu(self.norm2(h))) + return x + h + + +class AnimaLLLiteASPP(nn.Module): + def __init__(self, channels, dilations, device=None, dtype=None, operations=None): + super().__init__() + branches = [] + for dilation in dilations: + if dilation == 1: + conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype) + else: + conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype) + branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU())) + self.branches = nn.ModuleList(branches) + self.global_pool = nn.AdaptiveAvgPool2d(1) + self.global_conv = nn.Sequential( + operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype), + _group_norm(channels, device=device, dtype=dtype, operations=operations), + nn.SiLU(), + ) + self.proj = nn.Sequential( + operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype), + _group_norm(channels, device=device, dtype=dtype, operations=operations), + nn.SiLU(), + ) + + def forward(self, x): + height, width = x.shape[-2:] + outputs = [branch(x) for branch in self.branches] + pooled = self.global_conv(self.global_pool(x)) + outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False)) + return self.proj(torch.cat(outputs, dim=1)) + + +class AnimaLLLiteConditioning(nn.Module): + def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None): + super().__init__() + half_dim = cond_dim // 2 + self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype) + self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations) + self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype) + self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations) + self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype) + self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations) + self.resblocks = nn.ModuleList([ + AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations) + for _ in range(cond_resblocks) + ]) + self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None + self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype) + self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype) + + def forward(self, x): + x = F.silu(self.norm1(self.conv1(x))) + x = F.silu(self.norm2(self.conv2(x))) + x = F.silu(self.norm3(self.conv3(x))) + for block in self.resblocks: + x = block(x) + if self.aspp is not None: + x = self.aspp(x) + x = self.proj(x).flatten(2).transpose(1, 2).contiguous() + return self.out_norm(x) + + +class AnimaLLLiteModule(nn.Module): + def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None): + super().__init__() + self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype) + self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype) + self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype) + self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype) + self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False) + + def forward(self, x, cond_emb, strength): + original_shape = x.shape + if x.ndim == 5: + x = x.flatten(1, 3) + + if x.shape[0] != cond_emb.shape[0]: + if x.shape[0] % cond_emb.shape[0] != 0: + raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}") + cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1) + if x.shape[1] != cond_emb.shape[1]: + raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}") + + cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb) + hidden = F.silu(self.down(x)) + gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1) + hidden = self.mid(torch.cat((cond_local, hidden), dim=-1)) + hidden = F.silu(hidden * (1 + gamma) + beta) + x = x + self.up(hidden) * strength + + if len(original_shape) == 5: + x = x.reshape(original_shape) + return x + + +class AnimaLLLite(nn.Module): + def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None): + super().__init__() + metadata = metadata or {} + version = metadata.get("lllite.version", "2") + if version != "2": + raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported") + + module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")}) + if not module_names: + raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules") + + cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1] + cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0] + cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0] + resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")} + cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0 + use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict) + dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8") + aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else () + + self.cond_in_channels = cond_in_channels + self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true" + self.lllite_conditioning1 = AnimaLLLiteConditioning( + cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, + device=device, dtype=dtype, operations=operations, + ) + + self.module_names = set() + self.block_count = 0 + self.model_dim = None + for name in module_names: + match = MODULE_PATTERN.fullmatch(name) + if match is None: + raise ValueError(f"Unsupported Anima LLLite module name: {name}") + down_shape = state_dict[f"{name}.down.weight"].shape + mlp_dim, in_dim = down_shape + module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1] + if module_cond_dim != cond_emb_dim: + raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}") + if self.model_dim is None: + self.model_dim = in_dim + elif self.model_dim != in_dim: + raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}") + self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations)) + self.module_names.add(name) + self.block_count = max(self.block_count, int(match.group(1)) + 1) + + def encode_conditioning(self, image): + return self.lllite_conditioning1(image) + + def apply(self, x, cond_emb, block_index, target, strength): + name = f"lllite_dit_blocks_{block_index}_{target}" + if name not in self.module_names: + return x + return self.get_submodule(name)(x, cond_emb, strength) + + +class AnimaLLLitePatch: + def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end): + self.model_patch = model_patch + self.image = image + self.mask = mask + self.strength = strength + self.sigma_start = sigma_start + self.sigma_end = sigma_end + + def __call__(self, args): + x = args["x"] + transformer_options = args["transformer_options"] + if self.strength == 0.0: + return args + sigmas = transformer_options.get("sigmas") + if sigmas is not None: + sigma = float(sigmas.max().item()) + if not self.sigma_end <= sigma <= self.sigma_start: + return args + if x.shape[2] != 1: + raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}") + + target_height = x.shape[-2] * 8 + target_width = x.shape[-1] * 8 + image = comfy.utils.common_upscale( + self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center" + ).clamp(0.0, 1.0) + image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0 + + if self.model_patch.model.cond_in_channels == 4: + mask = self.mask + if mask.ndim == 3: + mask = mask.unsqueeze(1) + if mask.ndim != 4 or mask.shape[1] != 1: + raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}") + mask = comfy.utils.common_upscale( + mask.float(), target_width, target_height, "nearest-exact", crop="center" + ) + if mask.shape[0] != image.shape[0]: + if image.shape[0] % mask.shape[0] != 0: + raise ValueError( + f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}" + ) + mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1) + mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype) + if self.model_patch.model.inpaint_masked_input: + image = image * (mask < 0.5).to(image.dtype) + image = torch.cat((image, mask * 2.0 - 1.0), dim=1) + + cond_emb = self.model_patch.model.encode_conditioning(image) + transformer_options["model_patch_data"][self] = cond_emb + return args + + def to(self, device_or_dtype): + return self + + def models(self): + return [self.model_patch] + + +class AnimaLLLiteAttentionPatch: + def __init__(self, patch, targets): + self.patch = patch + self.targets = targets + + def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None): + cond_emb = extra_options["model_patch_data"].get(self.patch) + if cond_emb is None: + return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask} + + block_index = extra_options["block_index"] + values = {"q": q, "k": k, "v": v} + for value_name, target in self.targets.items(): + values[value_name] = self.patch.model_patch.model.apply( + values[value_name], cond_emb, block_index, target, self.patch.strength + ) + + return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask} + + +class AnimaLLLiteMLPPatch: + def __init__(self, patch): + self.patch = patch + + def __call__(self, args): + cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch) + if cond_emb is None: + return args + args["x"] = self.patch.model_patch.model.apply( + args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength + ) + return args diff --git a/comfy/ldm/audio/dit.py b/comfy/ldm/audio/dit.py index c28be5b49..b0759a240 100644 --- a/comfy/ldm/audio/dit.py +++ b/comfy/ldm/audio/dit.py @@ -425,19 +425,16 @@ class Attention(nn.Module): if n == 1 and causal: causal = False - if h != kv_h: - # Repeat interleave kv_heads to match q_heads - heads_per_kv_head = h // kv_h - k, v = map(lambda t: t.repeat_interleave(heads_per_kv_head, dim = 1), (k, v)) + gqa_kwargs = {"enable_gqa": True} if h != kv_h else {} if self.differential: q, q_diff = q.unbind(dim=1) k, k_diff = k.unbind(dim=1) - out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) - out_diff = optimized_attention(q_diff, k_diff, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) + out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) + out_diff = optimized_attention(q_diff, k_diff, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) out = out - out_diff else: - out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options) + out = optimized_attention(q, k, v, h, skip_reshape=True, low_precision_attention=False, transformer_options=transformer_options, **gqa_kwargs) out = self.to_out(out) diff --git a/comfy/ldm/boogu/model.py b/comfy/ldm/boogu/model.py index 966f3c583..ca88bdeb1 100644 --- a/comfy/ldm/boogu/model.py +++ b/comfy/ldm/boogu/model.py @@ -74,11 +74,8 @@ class BooguDoubleStreamProcessor(nn.Module): key = key.transpose(1, 2) value = value.transpose(1, 2) - if attn.kv_heads < attn.heads: - key = key.repeat_interleave(attn.heads // attn.kv_heads, dim=1) - value = value.repeat_interleave(attn.heads // attn.kv_heads, dim=1) - - hidden_states = optimized_attention_masked(query, key, value, attn.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options) + gqa_kwargs = {"enable_gqa": True} if attn.kv_heads < attn.heads else {} + hidden_states = optimized_attention_masked(query, key, value, attn.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options, **gqa_kwargs) # Split back to instruction/image, apply per-stream output projections, recombine. instruct_hidden_states = self.instruct_out(hidden_states[:, :L_instruct]) diff --git a/comfy/ldm/cosmos/predict2.py b/comfy/ldm/cosmos/predict2.py index aec874815..d391d50b1 100644 --- a/comfy/ldm/cosmos/predict2.py +++ b/comfy/ldm/cosmos/predict2.py @@ -14,6 +14,7 @@ from torchvision import transforms import comfy.patcher_extension from comfy.ldm.modules.attention import optimized_attention import comfy.ldm.common_dit +import comfy.ops import comfy.quant_ops @@ -148,11 +149,29 @@ class Attention(nn.Module): x: torch.Tensor, context: Optional[torch.Tensor] = None, rope_emb: Optional[torch.Tensor] = None, + transformer_options: Optional[dict] = {}, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - q = self.q_proj(x) context = x if context is None else context - k = self.k_proj(context) - v = self.v_proj(context) + q_input = x + k_input = context + v_input = context + + transformer_patches = transformer_options.get("patches", {}) + patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch" + if patch_name in transformer_patches: + extra_options = transformer_options.copy() + extra_options["n_heads"] = self.n_heads + extra_options["dim_head"] = self.head_dim + for patch in transformer_patches[patch_name]: + out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options) + q_input = out.get("q", q_input) + k_input = out.get("k", k_input) + v_input = out.get("v", v_input) + rope_emb = out.get("pe", rope_emb) + + q = self.q_proj(q_input) + k = self.k_proj(k_input) + v = self.v_proj(v_input) q, k, v = map( lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim), (q, k, v), @@ -161,11 +180,16 @@ class Attention(nn.Module): def apply_norm_and_rotary_pos_emb( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor] ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - q = self.q_norm(q) - k = self.k_norm(k) v = self.v_norm(v) if self.is_selfattn and rope_emb is not None: # only apply to self-attention! - q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rope_emb) + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, q, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, k, offloadable=True) + q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, rope_emb, q_scale, k_scale, self.q_norm.eps) + comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream) + else: + q = self.q_norm(q) + k = self.k_norm(k) return q, k, v q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb) @@ -188,7 +212,7 @@ class Attention(nn.Module): x (Tensor): The query tensor of shape [B, Mq, K] context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None """ - q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb) + q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options) return self.compute_attention(q, k, v, transformer_options=transformer_options) @@ -555,8 +579,14 @@ class Block(nn.Module): self.layer_norm_mlp, scale_mlp_B_T_1_1_D, shift_mlp_B_T_1_1_D, - ) - result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype)) + ).to(compute_dtype) + patches = transformer_options.get("patches", {}) + if "mlp_patch" in patches: + args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options} + for patch in patches["mlp_patch"]: + args = patch(args) + normalized_x_B_T_H_W_D = args["x"] + result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D) x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype)) return x_B_T_H_W_D @@ -863,11 +893,22 @@ class MiniTrainDIT(nn.Module): x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape ), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}" + transformer_options = kwargs.get("transformer_options", {}) + patches = transformer_options.get("patches", {}) + if "post_input" in patches: + transformer_options = transformer_options.copy() + transformer_options["model_patch_data"] = {} + + if "post_input" in patches: + for patch in patches["post_input"]: + out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options}) + x_B_T_H_W_D = out["img"] + block_kwargs = { "rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0), "adaln_lora_B_T_3D": adaln_lora_B_T_3D, "extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D, - "transformer_options": kwargs.get("transformer_options", {}), + "transformer_options": transformer_options, } # The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream @@ -877,7 +918,8 @@ class MiniTrainDIT(nn.Module): if x_B_T_H_W_D.dtype == torch.float16: x_B_T_H_W_D = x_B_T_H_W_D.float() - for block in self.blocks: + for block_index, block in enumerate(self.blocks): + transformer_options["block_index"] = block_index x_B_T_H_W_D = block( x_B_T_H_W_D, t_embedding_B_T_D, diff --git a/comfy/ldm/ernie/model.py b/comfy/ldm/ernie/model.py index f158ca1d2..88a3775d0 100644 --- a/comfy/ldm/ernie/model.py +++ b/comfy/ldm/ernie/model.py @@ -5,6 +5,7 @@ import torch.nn.functional as F from comfy.ldm.modules.attention import optimized_attention import comfy.model_management +import comfy.ops import comfy.quant_ops def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: @@ -111,11 +112,17 @@ class ErnieImageAttention(nn.Module): query = q_flat.view(B, S, self.heads, self.head_dim) key = k_flat.view(B, S, self.heads, self.head_dim) - query = self.norm_q(query) - key = self.norm_k(key) - - if image_rotary_emb is not None: - query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb) + if image_rotary_emb is not None and not comfy.model_management.in_training: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, query, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, key, offloadable=True) + query, key = comfy.quant_ops.ck.rms_rope_split_half(query, key, image_rotary_emb, q_scale, k_scale, self.norm_q.eps) + comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream) + else: + query = self.norm_q(query) + key = self.norm_k(key) + if image_rotary_emb is not None: + query, key = comfy.quant_ops.ck.apply_rope_split_half(query, key, image_rotary_emb) q_flat = query.reshape(B, S, -1) k_flat = key.reshape(B, S, -1) diff --git a/comfy/ldm/hidream_o1/attention.py b/comfy/ldm/hidream_o1/attention.py index 1b68f1771..afb2be9b8 100644 --- a/comfy/ldm/hidream_o1/attention.py +++ b/comfy/ldm/hidream_o1/attention.py @@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None): The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes. """ - def two_pass_attention(q, k, v, heads, **kwargs): + def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs): B, H, T, D = q.shape if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) elif ar_len >= T: - out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True) + out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa) elif ar_len <= 0: - out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options) + out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa) else: out_ar = comfy.ops.scaled_dot_product_attention( q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len], - attn_mask=None, dropout_p=0.0, is_causal=True, + attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa, ) out_gen = optimized_attention( q[:, :, ar_len:], k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, - transformer_options=transformer_options, + transformer_options=transformer_options, enable_gqa=enable_gqa, ) out = torch.cat([out_ar, out_gen], dim=2) diff --git a/comfy/ldm/ideogram4/model.py b/comfy/ldm/ideogram4/model.py index 4ea5b8aaf..12e1a14fb 100644 --- a/comfy/ldm/ideogram4/model.py +++ b/comfy/ldm/ideogram4/model.py @@ -12,10 +12,13 @@ import torch import torch.nn as nn import torch.nn.functional as F +import comfy.model_management +import comfy.ops import comfy.patcher_extension +import comfy.quant_ops from comfy.ldm.lumina.model import FeedForward from comfy.ldm.modules.attention import optimized_attention_masked -from comfy.text_encoders.llama import apply_rope, precompute_freqs_cis +from comfy.text_encoders.llama import precompute_freqs_cis # Per-token role indicators SEQUENCE_PADDING_INDICATOR = -1 @@ -25,6 +28,22 @@ LLM_TOKEN_INDICATOR = 3 IMAGE_POSITION_OFFSET = 65536 +def _split_half_rope_matrix(freqs_cis): + cos, sin, neg_sin = freqs_cis + half_dim = sin.shape[-1] + matrix = torch.stack( + (cos[..., :half_dim], neg_sin, sin, cos[..., half_dim:]), dim=-1 + ) + return matrix.reshape(*matrix.shape[:-1], 2, 2).unsqueeze(2) + + +def _apply_rope_split_half1(x, freqs_cis): + x_dtype = x.dtype + x = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(freqs_cis.dtype) + output = freqs_cis[..., 0] * x[..., 0] + freqs_cis[..., 1] * x[..., 1] + return output.movedim(-1, -2).reshape(*x.shape[:-3], -1).to(x_dtype) + + class Ideogram4Attention(nn.Module): def __init__(self, hidden_size, num_heads, eps=1e-5, dtype=None, device=None, operations=None): super().__init__() @@ -42,16 +61,23 @@ class Ideogram4Attention(nn.Module): qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim) q, k, v = qkv.unbind(dim=2) - q = self.norm_q(q) - k = self.norm_k(k) + if comfy.model_management.in_training: + q = _apply_rope_split_half1(self.norm_q(q), freqs_cis) + k = _apply_rope_split_half1(self.norm_k(k), freqs_cis) + else: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.norm_q, q, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.norm_k, k, offloadable=True) + q, k = comfy.quant_ops.ck.rms_rope_split_half( + q, k, freqs_cis, q_scale, k_scale, self.norm_q.eps + ) + comfy.ops.uncast_bias_weight(self.norm_q, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.norm_k, k_scale, None, k_offload_stream) # (B, heads, L, head_dim) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) - q, k = apply_rope(q, k, freqs_cis) - out = optimized_attention_masked(q, k, v, self.num_heads, attn_mask, skip_reshape=True, transformer_options=transformer_options) return self.o(out) @@ -181,6 +207,7 @@ class Ideogram4Transformer(nn.Module): self.head_dim, position_ids[0].transpose(0, 1), self.rope_theta, rope_dims=self.mrope_section, interleaved_mrope=True, device=position_ids.device, ) + freqs_cis = _split_half_rope_matrix(freqs_cis) if attn_mask is not None and attn_mask.dtype == torch.bool: attn_mask = torch.zeros_like(attn_mask, dtype=h.dtype).masked_fill_(~attn_mask, -torch.finfo(h.dtype).max) diff --git a/comfy/ldm/joyimage/model.py b/comfy/ldm/joyimage/model.py new file mode 100644 index 000000000..9d6951e54 --- /dev/null +++ b/comfy/ldm/joyimage/model.py @@ -0,0 +1,454 @@ +# https://github.com/jdopensource/JoyAI-Image-Edit (Apache 2.0) +import math +from typing import Optional, Tuple + +import comfy_kitchen +import torch +import torch.nn as nn + +import comfy.ldm.common_dit +import comfy.ops +import comfy.patcher_extension +from comfy.ldm.lightricks.model import GELU_approx, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps +from comfy.ldm.modules.attention import optimized_attention + + +class JoyImageModulate(nn.Module): + def __init__(self, hidden_size: int, factor: int, dtype=None, device=None): + super().__init__() + self.factor = factor + self.modulate_table = nn.Parameter( + torch.empty(1, factor, hidden_size, dtype=dtype, device=device) + ) + + def forward(self, x: torch.Tensor) -> list: + if x.ndim != 3: + x = x.unsqueeze(1) + table = comfy.ops.cast_to_input(self.modulate_table, x) + return [o.squeeze(1) for o in (table + x).chunk(self.factor, dim=1)] + + +class JoyImageFeedForward(nn.Module): + def __init__( + self, + dim: int, + inner_dim: int, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.net = nn.ModuleList([ + GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations), + nn.Identity(), + operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device), + ]) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + for module in self.net: + x = module(x) + return x + + +class JoyImageAttention(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + eps: float = 1e-6, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.num_attention_heads = num_attention_heads + inner_dim = num_attention_heads * attention_head_dim + + self.img_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device) + self.img_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.img_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.img_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device) + + self.txt_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device) + self.txt_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.txt_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device) + self.txt_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device) + + def forward( + self, + img: torch.Tensor, + txt: torch.Tensor, + image_rotary_emb: torch.Tensor, + transformer_options=None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + heads = self.num_attention_heads + + img_q, img_k, img_v = self.img_attn_qkv(img).chunk(3, dim=-1) + txt_q, txt_k, txt_v = self.txt_attn_qkv(txt).chunk(3, dim=-1) + + img_q = img_q.unflatten(-1, (heads, -1)) + img_k = img_k.unflatten(-1, (heads, -1)) + img_v = img_v.unflatten(-1, (heads, -1)) + txt_q = txt_q.unflatten(-1, (heads, -1)) + txt_k = txt_k.unflatten(-1, (heads, -1)) + txt_v = txt_v.unflatten(-1, (heads, -1)) + + txt_q = self.txt_attn_q_norm(txt_q) + txt_k = self.txt_attn_k_norm(txt_k) + + img_q_scale, _, img_q_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_q_norm, img_q, offloadable=True) + img_k_scale, _, img_k_offload_stream = comfy.ops.cast_bias_weight(self.img_attn_k_norm, img_k, offloadable=True) + img_q, img_k = comfy_kitchen.rms_rope( + img_q, + img_k, + image_rotary_emb, + img_q_scale, + img_k_scale, + self.img_attn_q_norm.eps, + ) + comfy.ops.uncast_bias_weight(self.img_attn_q_norm, img_q_scale, None, img_q_offload_stream) + comfy.ops.uncast_bias_weight(self.img_attn_k_norm, img_k_scale, None, img_k_offload_stream) + + joint_q = torch.cat([img_q, txt_q], dim=1) + joint_k = torch.cat([img_k, txt_k], dim=1) + joint_v = torch.cat([img_v, txt_v], dim=1) + + joint_q = joint_q.flatten(2, 3) + joint_k = joint_k.flatten(2, 3) + joint_v = joint_v.flatten(2, 3) + + joint_out = optimized_attention(joint_q, joint_k, joint_v, heads=heads, transformer_options=transformer_options) + + seq_img = img.shape[1] + img_out = joint_out[:, :seq_img, :] + txt_out = joint_out[:, seq_img:, :] + + img_out = self.img_attn_proj(img_out) + txt_out = self.txt_attn_proj(txt_out) + return img_out, txt_out + + +class JoyImageTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + num_attention_heads: int, + attention_head_dim: int, + mlp_width_ratio: float = 4.0, + eps: float = 1e-6, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + mlp_hidden_dim = int(dim * mlp_width_ratio) + + self.img_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device) + self.img_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.img_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.img_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations) + + self.txt_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device) + self.txt_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.txt_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device) + self.txt_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations) + + self.attn = JoyImageAttention( + dim=dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + eps=eps, + dtype=dtype, + device=device, + operations=operations, + ) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + image_rotary_emb: torch.Tensor, + transformer_options=None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + ( + img_mod1_shift, + img_mod1_scale, + img_mod1_gate, + img_mod2_shift, + img_mod2_scale, + img_mod2_gate, + ) = self.img_mod(temb) + ( + txt_mod1_shift, + txt_mod1_scale, + txt_mod1_gate, + txt_mod2_shift, + txt_mod2_scale, + txt_mod2_gate, + ) = self.txt_mod(temb) + + img_normed = self.img_norm1(hidden_states) + txt_normed = self.txt_norm1(encoder_hidden_states) + img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1) + txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1) + + img_attn, txt_attn = self.attn(img_modulated, txt_modulated, image_rotary_emb, transformer_options=transformer_options) + + hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1) + encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1) + + img_ffn_normed = self.img_norm2(hidden_states) + txt_ffn_normed = self.txt_norm2(encoder_hidden_states) + img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1) + txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1) + hidden_states = hidden_states + self.img_mlp(img_ffn_input) * img_mod2_gate.unsqueeze(1) + encoder_hidden_states = encoder_hidden_states + self.txt_mlp(txt_ffn_input) * txt_mod2_gate.unsqueeze(1) + + return hidden_states, encoder_hidden_states + + +class JoyImageTimeTextImageEmbedding(nn.Module): + def __init__( + self, + dim: int, + time_freq_dim: int, + time_proj_dim: int, + text_embed_dim: int, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0) + self.time_embedder = TimestepEmbedding( + in_channels=time_freq_dim, + time_embed_dim=dim, + dtype=dtype, + device=device, + operations=operations, + ) + self.act_fn = nn.SiLU() + self.time_proj = operations.Linear(dim, time_proj_dim, bias=True, dtype=dtype, device=device) + self.text_embedder = PixArtAlphaTextProjection( + text_embed_dim, dim, act_fn="gelu_tanh", dtype=dtype, device=device, operations=operations, + ) + + def forward(self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor): + timestep = self.timesteps_proj(timestep) + temb = self.time_embedder(timestep.to(dtype=encoder_hidden_states.dtype)).type_as(encoder_hidden_states) + timestep_proj = self.time_proj(self.act_fn(temb)) + encoder_hidden_states = self.text_embedder(encoder_hidden_states) + return temb, timestep_proj, encoder_hidden_states + + +class JoyImageTransformer3DModel(nn.Module): + def __init__( + self, + patch_size: list = [1, 2, 2], + in_channels: int = 16, + out_channels: Optional[int] = None, + hidden_size: int = 3072, + num_attention_heads: int = 24, + text_dim: int = 4096, + mlp_width_ratio: float = 4.0, + num_layers: int = 20, + rope_dim_list: list = [16, 56, 56], + theta: int = 256, + image_model=None, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.dtype = dtype + self.out_channels = out_channels or in_channels + self.patch_size = list(patch_size) + self.rope_dim_list = list(rope_dim_list) + self.theta = theta + + attention_head_dim = hidden_size // num_attention_heads + + self.img_in = operations.Conv3d( + in_channels, + hidden_size, + kernel_size=tuple(self.patch_size), + stride=tuple(self.patch_size), + dtype=dtype, + device=device, + ) + + self.condition_embedder = JoyImageTimeTextImageEmbedding( + dim=hidden_size, + time_freq_dim=256, + time_proj_dim=hidden_size * 6, + text_embed_dim=text_dim, + dtype=dtype, + device=device, + operations=operations, + ) + + self.double_blocks = nn.ModuleList([ + JoyImageTransformerBlock( + dim=hidden_size, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + mlp_width_ratio=mlp_width_ratio, + dtype=dtype, + device=device, + operations=operations, + ) + for _ in range(num_layers) + ]) + + self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.proj_out = operations.Linear( + hidden_size, + self.out_channels * math.prod(self.patch_size), + bias=True, + dtype=dtype, + device=device, + ) + + def _get_rotary_pos_embed_for_range( + self, + start: Tuple[int, int, int], + stop: Tuple[int, int, int], + device=None, + ) -> torch.Tensor: + # 3D RoPE for the patch grid range [start, stop) over (t, h, w). Token order after + # reshape(-1) is (t, h, w), matching the img_in Conv3d flatten. + rope_dim_list = self.rope_dim_list + + grids = [torch.arange(start[i], stop[i], dtype=torch.float32, device=device) for i in range(3)] + mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0) + + angles_parts = [] + for i, dim in enumerate(rope_dim_list): + pos = mesh[i].reshape(-1) + freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device)[: (dim // 2)] / dim)) + angles_parts.append(torch.outer(pos, freqs)) + + angles = torch.cat(angles_parts, dim=1) + cos = angles.cos() + sin = angles.sin() + return torch.stack((cos, -sin, sin, cos), dim=-1).unflatten(-1, (2, 2)) + + def get_rotary_pos_embed_for_components( + self, + component_sizes, + device=None, + ) -> torch.Tensor: + # Per-component 3D RoPE. component_sizes is a list of (t, h, w) patch grid sizes in + # sequence order [target, ref0, ref1, ...]; h/w restart at 0 for each component while t + # continues from the running offset, giving every image its own temporal position band. + freqs_parts = [] + t_offset = 0 + for (t, h, w) in component_sizes: + freqs = self._get_rotary_pos_embed_for_range( + start=(t_offset, 0, 0), + stop=(t_offset + t, h, w), + device=device, + ) + freqs_parts.append(freqs) + t_offset += t + return torch.cat(freqs_parts, dim=0).unsqueeze(0).unsqueeze(2) + + def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor: + c = self.out_channels + pt, ph, pw = self.patch_size + x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c) + x = x.permute(0, 7, 1, 4, 2, 5, 3, 6) + return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw) + + def forward( + self, + hidden_states: torch.Tensor, + timestep: torch.Tensor, + context: torch.Tensor = None, + ref_latents=None, + control=None, + transformer_options=None, + **kwargs, + ) -> torch.Tensor: + transformer_options = {} if transformer_options is None else transformer_options.copy() + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(hidden_states, timestep, context, ref_latents, transformer_options, **kwargs) + + def _forward( + self, + hidden_states: torch.Tensor, + timestep: torch.Tensor, + context: torch.Tensor, + ref_latents=None, + transformer_options=None, + **kwargs, + ) -> torch.Tensor: + pt, ph, pw = self.patch_size + _, _, ot, oh, ow = hidden_states.shape + + components = [hidden_states, *(ref_latents or [])] + component_sizes = [] + img_tokens = [] + for comp in components: + comp = comfy.ldm.common_dit.pad_to_patch_size(comp, self.patch_size) + _, _, ct, ch, cw = comp.shape + component_sizes.append((ct // pt, ch // ph, cw // pw)) + tokens = self.img_in(comp).flatten(2).transpose(1, 2) # (B, n_i, D) + img_tokens.append(tokens) + + img = torch.cat(img_tokens, dim=1) + + _, vec, txt = self.condition_embedder(timestep, context) + vec = vec.unflatten(1, (6, -1)) + + image_rotary_emb = self.get_rotary_pos_embed_for_components( + component_sizes, + device=hidden_states.device, + ) + + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + transformer_options["total_blocks"] = len(self.double_blocks) + transformer_options["block_type"] = "double" + for i, block in enumerate(self.double_blocks): + transformer_options["block_index"] = i + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"], out["txt"] = block( + hidden_states=args["img"], + encoder_hidden_states=args["txt"], + temb=args["vec"], + image_rotary_emb=args["pe"], + transformer_options=args.get("transformer_options"), + ) + return out + + out = blocks_replace[("double_block", i)]({"img": img, + "txt": txt, + "vec": vec, + "pe": image_rotary_emb, + "transformer_options": transformer_options}, + {"original_block": block_wrap}) + txt = out["txt"] + img = out["img"] + else: + img, txt = block( + hidden_states=img, + encoder_hidden_states=txt, + temb=vec, + image_rotary_emb=image_rotary_emb, + transformer_options=transformer_options, + ) + + tt, th, tw = component_sizes[0] + target_tokens = tt * th * tw + img = img[:, :target_tokens, :] + img = self.proj_out(self.norm_out(img)) + img = self.unpatchify(img, tt, th, tw) + return img[:, :, :ot, :oh, :ow] diff --git a/comfy/ldm/krea2/model.py b/comfy/ldm/krea2/model.py new file mode 100644 index 000000000..8001812d7 --- /dev/null +++ b/comfy/ldm/krea2/model.py @@ -0,0 +1,391 @@ +"""Krea 2 (K2) — single-stream MMDiT. + +Text tokens produced by a Qwen3-VL-4B 12-layer ``txtfusion`` adapter and patchified image tokens are +concatenated into one sequence and run through ``layers`` shared transformer blocks with +AdaLN-single modulation, GQA + per-head QK-norm + sigmoid-gated attention, SwiGLU MLP, and 3-axis RoPE. +""" + +from typing import Optional + +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange + +import comfy.model_management +import comfy.patcher_extension +import comfy.ldm.common_dit +import comfy.utils +from comfy.ldm.flux.layers import EmbedND, timestep_embedding +from comfy.ldm.flux.math import apply_rope +from comfy.ldm.modules.attention import optimized_attention_masked + + +class RMSNorm(nn.Module): + """RMSNorm with the reference ``(1 + scale)`` weight convention (scale stored zero-centered).""" + + def __init__(self, features: int, eps: float = 1e-5, device=None, dtype=None, operations=None): + super().__init__() + self.eps = eps + self.scale = nn.Parameter(torch.empty(features, device=device, dtype=dtype)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + weight = comfy.model_management.cast_to(self.scale, dtype=torch.float32, device=x.device) + 1.0 + return F.rms_norm(x.float(), (x.shape[-1],), weight=weight, eps=self.eps).to(dtype) + + +class QKNorm(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.qnorm = RMSNorm(dim, device=device, dtype=dtype, operations=operations) + self.knorm = RMSNorm(dim, device=device, dtype=dtype, operations=operations) + + def forward(self, q, k): + return self.qnorm(q), self.knorm(k) + + +class SwiGLU(nn.Module): + def __init__(self, features: int, multiplier: int, bias: bool = False, multiple: int = 128, + device=None, dtype=None, operations=None): + super().__init__() + mlpdim = int(2 * features / 3) * multiplier + mlpdim = multiple * ((mlpdim + multiple - 1) // multiple) + self.gate = operations.Linear(features, mlpdim, bias=bias, device=device, dtype=dtype) + self.up = operations.Linear(features, mlpdim, bias=bias, device=device, dtype=dtype) + self.down = operations.Linear(mlpdim, features, bias=bias, device=device, dtype=dtype) + + def forward(self, x): + return self.down(F.silu(self.gate(x)).mul_(self.up(x))) + + +class Attention(nn.Module): + def __init__(self, dim: int, heads: int, kvheads: Optional[int] = None, bias: bool = False, + device=None, dtype=None, operations=None): + super().__init__() + self.heads = heads + self.kvheads = kvheads if kvheads is not None else heads + self.headdim = dim // self.heads + self.wq = operations.Linear(dim, self.headdim * self.heads, bias=bias, device=device, dtype=dtype) + self.wk = operations.Linear(dim, self.headdim * self.kvheads, bias=bias, device=device, dtype=dtype) + self.wv = operations.Linear(dim, self.headdim * self.kvheads, bias=bias, device=device, dtype=dtype) + self.gate = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) + self.qknorm = QKNorm(self.headdim, device=device, dtype=dtype, operations=operations) + self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype) + + def forward(self, x, freqs=None, mask=None, transformer_options={}): + transformer_patches = transformer_options.get("patches", {}) + extra_options = transformer_options.copy() + q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x) + q = rearrange(q, "B L (H D) -> B H L D", H=self.heads) + k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads) + v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads) + q, k = self.qknorm(q, k) + + if "block_index" in transformer_options and "attn1_patch" in transformer_patches: + for p in transformer_patches["attn1_patch"]: + out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options) + q, k, v = out.get("q", q), out.get("k", k), out.get("v", v) + freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask) + + if freqs is not None: + q, k = apply_rope(q, k, freqs) + if self.kvheads != self.heads: + rep = self.heads // self.kvheads + k = k.repeat_interleave(rep, dim=1) + v = v.repeat_interleave(rep, dim=1) + out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True, + transformer_options=transformer_options) + + if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches: + for p in transformer_patches["attn1_output_patch"]: + out = p(out, extra_options) + + return self.wo(out * F.sigmoid(gate)) + + +class SimpleModulation(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.lin = nn.Parameter(torch.empty(2, dim, device=device, dtype=dtype)) + + def forward(self, vec): + out = vec + comfy.model_management.cast_to(self.lin, dtype=vec.dtype, device=vec.device).unsqueeze(0) + scale, shift = out.chunk(2, dim=1) + return scale, shift + + +class DoubleSharedModulation(nn.Module): + def __init__(self, dim: int, device=None, dtype=None, operations=None): + super().__init__() + self.lin = nn.Parameter(torch.empty(6 * dim, device=device, dtype=dtype)) + + def forward(self, vec): + out = vec + comfy.model_management.cast_to(self.lin, dtype=vec.dtype, device=vec.device) + return out.chunk(6, dim=-1) + + +class TextFusionBlock(nn.Module): + def __init__(self, features, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.prenorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.postnorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) + self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) + + def forward(self, x, mask=None, transformer_options={}): + x = x + self.attn(self.prenorm(x), mask=mask, transformer_options=transformer_options) + x = x + self.mlp(self.postnorm(x)) + return x + + +class TextFusionTransformer(nn.Module): + def __init__(self, num_txt_layers, txt_dim, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.layerwise_blocks = nn.ModuleList([ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(2) + ]) + self.projector = operations.Linear(num_txt_layers, 1, bias=False, device=device, dtype=dtype) + self.refiner_blocks = nn.ModuleList([ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(2) + ]) + + def forward(self, x, mask=None, transformer_options={}): + b, l, n, d = x.shape + x = x.reshape(b * l, n, d) + for block in self.layerwise_blocks: + x = block(x.contiguous(), mask=None, transformer_options=transformer_options) + x = rearrange(x, "(b l) n d -> b l d n", b=b, l=l) + x = self.projector(x).squeeze(-1) + for block in self.refiner_blocks: + x = block(x, mask=mask, transformer_options=transformer_options) + return x + + +class SingleStreamBlock(nn.Module): + def __init__(self, features, heads, multiplier, bias=False, kvheads=None, device=None, dtype=None, operations=None): + super().__init__() + self.mod = DoubleSharedModulation(features, device=device, dtype=dtype, operations=operations) + self.prenorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.postnorm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations) + self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations) + + def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}): + prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec) + if timestep_zero_index is not None: + bs = x.shape[0] + ref_prescale = prescale[bs:] + ref_preshift = preshift[bs:] + ref_pregate = pregate[bs:] + ref_postscale = postscale[bs:] + ref_postshift = postshift[bs:] + ref_postgate = postgate[bs:] + prescale = prescale[:bs] + preshift = preshift[:bs] + pregate = pregate[:bs] + postscale = postscale[:bs] + postshift = postshift[:bs] + postgate = postgate[:bs] + + pre = self.prenorm(x) + pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift) + pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift) + attn = self.attn(pre, freqs, mask, transformer_options=transformer_options) + del pre + attn[:, :timestep_zero_index].mul_(pregate) + attn[:, timestep_zero_index:].mul_(ref_pregate) + x = x + attn + del attn + + post = self.postnorm(x) + post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift) + post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift) + mlp = self.mlp(post) + del post + mlp[:, :timestep_zero_index].mul_(postgate) + mlp[:, timestep_zero_index:].mul_(ref_postgate) + x = x + mlp + del mlp + return x + + x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options) + x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift) + return x + + +class LastLayer(nn.Module): + def __init__(self, features, patch, channels, device=None, dtype=None, operations=None): + super().__init__() + self.norm = RMSNorm(features, device=device, dtype=dtype, operations=operations) + self.linear = operations.Linear(features, patch * patch * channels, bias=True, device=device, dtype=dtype) + self.modulation = SimpleModulation(features, device=device, dtype=dtype, operations=operations) + + def forward(self, x, tvec): + scale, shift = self.modulation(tvec) + x = (1 + scale) * self.norm(x) + shift + return self.linear(x) + + +class SingleStreamDiT(nn.Module): + def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4, + layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12, + txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None, + device=None, dtype=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.patch = patch + self.channels = channels + self.tdim = tdim + self.heads = heads + self.txtdim = txtdim + self.txtlayers = txtlayers + self.default_ref_method = default_ref_method + + headdim = features // heads + axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)] + assert sum(axes) == headdim, f"axes {axes} sum != headdim {headdim}" + self.pe_embedder = EmbedND(dim=headdim, theta=int(theta), axes_dim=axes) + + self.first = operations.Linear(channels * patch ** 2, features, bias=True, device=device, dtype=dtype) + self.blocks = nn.ModuleList([ + SingleStreamBlock(features, heads, multiplier, bias, kvheads, device=device, dtype=dtype, operations=operations) + for _ in range(layers) + ]) + self.tmlp = nn.Sequential( + operations.Linear(tdim, features, device=device, dtype=dtype), + nn.GELU(approximate="tanh"), + operations.Linear(features, features, device=device, dtype=dtype), + ) + self.txtfusion = TextFusionTransformer(txtlayers, txtdim, txtheads, multiplier, bias, txtkvheads, + device=device, dtype=dtype, operations=operations) + self.txtmlp = nn.Sequential( + RMSNorm(txtdim, device=device, dtype=dtype, operations=operations), + operations.Linear(txtdim, features, device=device, dtype=dtype), + nn.GELU(approximate="tanh"), + operations.Linear(features, features, device=device, dtype=dtype), + ) + self.last = LastLayer(features, patch, channels, device=device, dtype=dtype, operations=operations) + self.tproj = nn.Sequential( + nn.GELU(approximate="tanh"), + operations.Linear(features, features * 6, device=device, dtype=dtype), + ) + + def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options), + ).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs) + + def process_img(self, x, index=0): + patch = self.patch + x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch)) + h, w = x.shape[-2] // patch, x.shape[-1] // patch + img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch) + + img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32) + img_ids[..., 0] = index + img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None] + img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :] + return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w + + def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): + transformer_options = transformer_options.copy() + temporal = x.ndim == 5 + if temporal: + b5, c5, t5, h5, w5 = x.shape + x = x.reshape(b5 * t5, c5, h5, w5) + bs, _, h_orig, w_orig = x.shape + patch = self.patch + + # context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim). + context = self._unpack_context(context) + + img, imgpos, h_, w_ = self.process_img(x) + img_tokens = img.shape[1] + timestep_zero_index = None + ref_method = kwargs.get("ref_latents_method", self.default_ref_method) + if ref_method is not None and ref_latents is not None and len(ref_latents) > 0: + ref_tokens = [] + ref_pos = [] + ref_num_tokens = [] + for index, ref in enumerate(ref_latents, 1): + if ref.ndim == 5: + rb, rc, rt, rh5, rw5 = ref.shape + ref = ref.reshape(rb * rt, rc, rh5, rw5) + ref = comfy.utils.repeat_to_batch_size(ref, bs) + kontext, kontext_ids, _, _ = self.process_img(ref, index=index) + ref_tokens.append(kontext) + ref_pos.append(kontext_ids) + ref_num_tokens.append(kontext.shape[1]) + img = torch.cat([img] + ref_tokens, dim=1) + imgpos = torch.cat([imgpos] + ref_pos, dim=1) + del ref_tokens, ref_pos + if ref_method == "index_timestep_zero": + timestep_zero_index = img_tokens + transformer_options["reference_image_num_tokens"] = ref_num_tokens + + img = self.first(img) + + t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype)) + tvec = self.tproj(t) + if timestep_zero_index is not None: + t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype)) + tvec = torch.cat((tvec, self.tproj(t0)), dim=0) + + context = self.txtfusion(context, mask=None, transformer_options=transformer_options) + context = self.txtmlp(context) + + txtlen = context.shape[1] + device = context.device + txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32) + + patches = transformer_options.get("patches", {}) + if "post_input" in patches: + for p in patches["post_input"]: + out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options}) + img, context = out["img"], out["txt"] + imgpos, txtpos = out["img_ids"], out["txt_ids"] + + combined = torch.cat((context, img), dim=1) + del context, img + if timestep_zero_index is not None: + timestep_zero_index += txtlen + + # Position ids: text at 0, image at (0, h_idx, w_idx). + pos = torch.cat((txtpos, imgpos), dim=1) + del txtpos, imgpos + + freqs = self.pe_embedder(pos) + del pos + + transformer_options["total_blocks"] = len(self.blocks) + transformer_options["block_type"] = "single" + transformer_options["img_slice"] = [txtlen, combined.shape[1]] + for i, block in enumerate(self.blocks): + transformer_options["block_index"] = i + combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options) + + final = self.last(combined, t) + del combined + out = final[:, txtlen:txtlen + img_tokens, :] + out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)", + h=h_, w=w_, ph=patch, pw=patch, c=self.channels) + out = out[:, :, :h_orig, :w_orig] # crop padding back off + if temporal: + out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2) + return out + + def _unpack_context(self, context): + # context: (B, seq, txtlayers*txtdim) -> (B, seq, txtlayers, txtdim). + b, seq, fused = context.shape + if fused != self.txtlayers * self.txtdim: + raise ValueError( + f"Krea2 expects conditioning with {self.txtlayers}x{self.txtdim}={self.txtlayers * self.txtdim} " + f"features (a {self.txtlayers}-layer Qwen3-VL stack) but got {fused}. " + f"Load the text encoder with CLIPLoader type 'krea2'." + ) + return context.reshape(b, seq, self.txtlayers, self.txtdim) diff --git a/comfy/ldm/lightricks/embeddings_connector.py b/comfy/ldm/lightricks/embeddings_connector.py index 2811080be..1a6ddcc8d 100644 --- a/comfy/ldm/lightricks/embeddings_connector.py +++ b/comfy/ldm/lightricks/embeddings_connector.py @@ -6,9 +6,8 @@ import torch from comfy.ldm.lightricks.model import ( CrossAttention, FeedForward, + freqs_cis_matrix, generate_freq_grid_np, - interleaved_freqs_cis, - split_freqs_cis, ) from torch import nn @@ -244,12 +243,15 @@ class Embeddings1DConnector(nn.Module): expected_freqs = dim // 2 current_freqs = freqs.shape[-1] pad_size = expected_freqs - current_freqs - cos_freq, sin_freq = split_freqs_cis( - freqs, pad_size, self.num_attention_heads - ) else: - cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) - return cos_freq.to(dtype=out_dtype), sin_freq.to(dtype=out_dtype), self.split_rope + pad_size = dim % n_elem + return freqs_cis_matrix( + freqs, + pad_size, + self.split_rope, + self.num_attention_heads, + out_dtype, + ) def forward( self, diff --git a/comfy/ldm/lightricks/latent_upsampler.py b/comfy/ldm/lightricks/latent_upsampler.py index 78ed7653f..6a4beb1bf 100644 --- a/comfy/ldm/lightricks/latent_upsampler.py +++ b/comfy/ldm/lightricks/latent_upsampler.py @@ -97,11 +97,11 @@ class SpatialRationalResampler(nn.Module): For dims==3, work per-frame for spatial scaling (temporal axis untouched). """ - def __init__(self, mid_channels: int, scale: float): + def __init__(self, mid_channels: int, scale: float, operations): super().__init__() self.scale = float(scale) self.num, self.den = _rational_for_scale(self.scale) - self.conv = nn.Conv2d( + self.conv = operations.Conv2d( mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1 ) self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num)) @@ -119,18 +119,18 @@ class SpatialRationalResampler(nn.Module): class ResBlock(nn.Module): def __init__( - self, channels: int, mid_channels: Optional[int] = None, dims: int = 3 + self, channels: int, operations, mid_channels: Optional[int] = None, dims: int = 3 ): super().__init__() if mid_channels is None: mid_channels = channels - Conv = nn.Conv2d if dims == 2 else nn.Conv3d + Conv = operations.Conv2d if dims == 2 else operations.Conv3d self.conv1 = Conv(channels, mid_channels, kernel_size=3, padding=1) - self.norm1 = nn.GroupNorm(32, mid_channels) + self.norm1 = operations.GroupNorm(32, mid_channels) self.conv2 = Conv(mid_channels, channels, kernel_size=3, padding=1) - self.norm2 = nn.GroupNorm(32, channels) + self.norm2 = operations.GroupNorm(32, channels) self.activation = nn.SiLU() def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -159,6 +159,7 @@ class LatentUpsampler(nn.Module): def __init__( self, + operations, in_channels: int = 128, mid_channels: int = 512, num_blocks_per_stage: int = 4, @@ -179,34 +180,34 @@ class LatentUpsampler(nn.Module): self.spatial_scale = float(spatial_scale) self.rational_resampler = rational_resampler - Conv = nn.Conv2d if dims == 2 else nn.Conv3d + Conv = operations.Conv2d if dims == 2 else operations.Conv3d self.initial_conv = Conv(in_channels, mid_channels, kernel_size=3, padding=1) - self.initial_norm = nn.GroupNorm(32, mid_channels) + self.initial_norm = operations.GroupNorm(32, mid_channels) self.initial_activation = nn.SiLU() self.res_blocks = nn.ModuleList( - [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] + [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)] ) if spatial_upsample and temporal_upsample: self.upsampler = nn.Sequential( - nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), + operations.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(3), ) elif spatial_upsample: if rational_resampler: self.upsampler = SpatialRationalResampler( - mid_channels=mid_channels, scale=self.spatial_scale + mid_channels=mid_channels, scale=self.spatial_scale, operations=operations ) else: self.upsampler = nn.Sequential( - nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), + operations.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(2), ) elif temporal_upsample: self.upsampler = nn.Sequential( - nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), + operations.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1), PixelShuffleND(1), ) else: @@ -215,11 +216,14 @@ class LatentUpsampler(nn.Module): ) self.post_upsample_res_blocks = nn.ModuleList( - [ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)] + [ResBlock(mid_channels, dims=dims, operations=operations) for _ in range(num_blocks_per_stage)] ) self.final_conv = Conv(mid_channels, in_channels, kernel_size=3, padding=1) + def get_dtype(self): + return getattr(self.initial_conv, "weight_comfy_model_dtype", self.initial_conv.weight.dtype) + def forward(self, latent: torch.Tensor) -> torch.Tensor: b, c, f, h, w = latent.shape @@ -266,7 +270,7 @@ class LatentUpsampler(nn.Module): return x @classmethod - def from_config(cls, config): + def from_config(cls, config, operations): return cls( in_channels=config.get("in_channels", 4), mid_channels=config.get("mid_channels", 128), @@ -276,6 +280,7 @@ class LatentUpsampler(nn.Module): temporal_upsample=config.get("temporal_upsample", False), spatial_scale=config.get("spatial_scale", 2.0), rational_resampler=config.get("rational_resampler", False), + operations=operations, ) def config(self): diff --git a/comfy/ldm/lightricks/model.py b/comfy/ldm/lightricks/model.py index 9953b6679..f9de3a38e 100644 --- a/comfy/ldm/lightricks/model.py +++ b/comfy/ldm/lightricks/model.py @@ -12,6 +12,8 @@ from torch import nn import comfy.patcher_extension import comfy.ldm.modules.attention import comfy.ldm.common_dit +import comfy.model_management +import comfy.quant_ops from .symmetric_patchifier import SymmetricPatchifier, latent_to_pixel_coords @@ -322,40 +324,42 @@ class FeedForward(nn.Module): return self.net(x) def apply_rotary_emb(input_tensor, freqs_cis): - cos_freqs, sin_freqs = freqs_cis[0], freqs_cis[1] - split_pe = freqs_cis[2] if len(freqs_cis) > 2 else False - return ( - apply_split_rotary_emb(input_tensor, cos_freqs, sin_freqs) - if split_pe else - apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs) + rotation_matrix, split_pe = freqs_cis + original_shape = input_tensor.shape + input_tensor = input_tensor.reshape( + input_tensor.shape[0], input_tensor.shape[1], rotation_matrix.shape[2], -1 ) -def apply_interleaved_rotary_emb(input_tensor, cos_freqs, sin_freqs): # TODO: remove duplicate funcs and pick the best/fastest one - t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2) - t1, t2 = t_dup.unbind(dim=-1) - t_dup = torch.stack((-t2, t1), dim=-1) - input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)") + if comfy.model_management.in_training: + if split_pe: + t = input_tensor.reshape(*input_tensor.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2) + else: + t = input_tensor.reshape(*input_tensor.shape[:-1], -1, 1, 2) + t = t.to(rotation_matrix.dtype) + output = rotation_matrix[..., 0] * t[..., 0] + rotation_matrix[..., 1] * t[..., 1] + if split_pe: + output = output.movedim(-1, -2) + output = output.reshape(input_tensor.shape).type_as(input_tensor) + elif split_pe: + output = comfy.quant_ops.ck.apply_rope_split_half1(input_tensor, rotation_matrix) + else: + output = comfy.quant_ops.ck.apply_rope1(input_tensor, rotation_matrix) + return output.reshape(original_shape) - out = input_tensor * cos_freqs + input_tensor_rot * sin_freqs +def apply_rotary_emb_qk(q, k, freqs_cis): + if comfy.model_management.in_training: + return apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis) - return out - -def apply_split_rotary_emb(input_tensor, cos, sin): - needs_reshape = False - if input_tensor.ndim != 4 and cos.ndim == 4: - B, H, T, _ = cos.shape - input_tensor = input_tensor.reshape(B, T, H, -1).swapaxes(1, 2) - needs_reshape = True - split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2) - first_half_input = split_input[..., :1, :] - second_half_input = split_input[..., 1:, :] - output = split_input * cos.unsqueeze(-2) - first_half_output = output[..., :1, :] - second_half_output = output[..., 1:, :] - first_half_output.addcmul_(-sin.unsqueeze(-2), second_half_input) - second_half_output.addcmul_(sin.unsqueeze(-2), first_half_input) - output = rearrange(output, "... d r -> ... (d r)") - return output.swapaxes(1, 2).reshape(B, T, -1) if needs_reshape else output + rotation_matrix, split_pe = freqs_cis + q_shape = q.shape + k_shape = k.shape + q = q.reshape(q.shape[0], q.shape[1], rotation_matrix.shape[2], -1) + k = k.reshape(k.shape[0], k.shape[1], rotation_matrix.shape[2], -1) + if split_pe: + q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rotation_matrix) + else: + q, k = comfy.quant_ops.ck.apply_rope(q, k, rotation_matrix) + return q.reshape(q_shape), k.reshape(k_shape) class GuideAttentionMask: @@ -461,9 +465,13 @@ class CrossAttention(nn.Module): 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: - q = apply_rotary_emb(q, pe) - k = apply_rotary_emb(k, pe if k_pe is None else k_pe) + 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) @@ -653,36 +661,23 @@ def generate_freqs(indices, indices_grid, max_pos, use_middle_indices_grid): ) return freqs -def interleaved_freqs_cis(freqs, pad_size): - cos_freq = freqs.cos().repeat_interleave(2, dim=-1) - sin_freq = freqs.sin().repeat_interleave(2, dim=-1) - if pad_size != 0: - cos_padding = torch.ones_like(cos_freq[:, :, : pad_size]) - sin_padding = torch.zeros_like(cos_freq[:, :, : pad_size]) - cos_freq = torch.cat([cos_padding, cos_freq], dim=-1) - sin_freq = torch.cat([sin_padding, sin_freq], dim=-1) - return cos_freq, sin_freq - -def split_freqs_cis(freqs, pad_size, num_attention_heads): - cos_freq = freqs.cos() - sin_freq = freqs.sin() - - if pad_size != 0: - cos_padding = torch.ones_like(cos_freq[:, :, :pad_size]) - sin_padding = torch.zeros_like(sin_freq[:, :, :pad_size]) - - cos_freq = torch.concatenate([cos_padding, cos_freq], axis=-1) - sin_freq = torch.concatenate([sin_padding, sin_freq], axis=-1) - - # Reshape freqs to be compatible with multi-head attention - B , T, half_HD = cos_freq.shape +def freqs_cis_matrix(freqs, pad_size, split_mode, num_attention_heads, out_dtype): + cos_freq = freqs.cos().to(out_dtype) + sin_freq = freqs.sin().to(out_dtype) + if pad_size: + matrix_pad_size = pad_size if split_mode else pad_size // 2 + cos_padding = torch.ones_like(cos_freq[:, :, :matrix_pad_size]) + sin_padding = torch.zeros_like(sin_freq[:, :, :matrix_pad_size]) + cos_freq = torch.cat((cos_padding, cos_freq), dim=-1) + sin_freq = torch.cat((sin_padding, sin_freq), dim=-1) + B, T, half_HD = cos_freq.shape cos_freq = cos_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) sin_freq = sin_freq.reshape(B, T, num_attention_heads, half_HD // num_attention_heads) - - cos_freq = torch.swapaxes(cos_freq, 1, 2) # (B,H,T,D//2) - sin_freq = torch.swapaxes(sin_freq, 1, 2) # (B,H,T,D//2) - return cos_freq, sin_freq + rotation_matrix = torch.stack( + (cos_freq, -sin_freq, sin_freq, cos_freq), dim=-1 + ) + return rotation_matrix.reshape(*rotation_matrix.shape[:-1], 2, 2), split_mode class LTXBaseModel(torch.nn.Module, ABC): """ @@ -885,12 +880,17 @@ class LTXBaseModel(torch.nn.Module, ABC): expected_freqs = dim // 2 current_freqs = freqs.shape[-1] pad_size = expected_freqs - current_freqs - cos_freq, sin_freq = split_freqs_cis(freqs, pad_size, num_attention_heads) else: # 2 because of cos and sin by 3 for (t, x, y), 1 for temporal only n_elem = 2 * indices_grid.shape[1] - cos_freq, sin_freq = interleaved_freqs_cis(freqs, dim % n_elem) - return cos_freq.to(out_dtype), sin_freq.to(out_dtype), split_mode + pad_size = dim % n_elem + return freqs_cis_matrix( + freqs, + pad_size, + split_mode, + num_attention_heads, + out_dtype, + ) def _prepare_positional_embeddings(self, pixel_coords, frame_rate, x_dtype): """Prepare positional embeddings.""" diff --git a/comfy/ldm/lightricks/vae/audio_vae.py b/comfy/ldm/lightricks/vae/audio_vae.py index dd5320c8f..b4a8c7524 100644 --- a/comfy/ldm/lightricks/vae/audio_vae.py +++ b/comfy/ldm/lightricks/vae/audio_vae.py @@ -185,7 +185,7 @@ class AudioVAE(torch.nn.Module): self.autoencoder.mel_bins, ) - def num_of_latents_from_frames(self, frames_number: int, frame_rate: int) -> int: + 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) def run_vocoder(self, mel_spec: torch.Tensor) -> torch.Tensor: diff --git a/comfy/ldm/lightricks/vae/causal_conv3d.py b/comfy/ldm/lightricks/vae/causal_conv3d.py index 7515f0d4e..bb1803f12 100644 --- a/comfy/ldm/lightricks/vae/causal_conv3d.py +++ b/comfy/ldm/lightricks/vae/causal_conv3d.py @@ -49,6 +49,12 @@ class CausalConv3d(nn.Module): ) self.temporal_cache_state={} + def _empty_output(self, x): + # empty (0 frame) outputs must still have the conv's output channels and spatial dims + h = (x.shape[3] + 2 * self.conv.padding[1] - self.conv.kernel_size[1]) // self.conv.stride[1] + 1 + w = (x.shape[4] + 2 * self.conv.padding[2] - self.conv.kernel_size[2]) // self.conv.stride[2] + 1 + return x.new_empty((x.shape[0], self.out_channels, 0, h, w)) + def forward(self, x, causal: bool = True): tid = threading.get_ident() @@ -58,7 +64,7 @@ class CausalConv3d(nn.Module): if not causal: padding_length = padding_length // 2 if x.shape[2] == 0: - return x + return self._empty_output(x) cached = x[:, :, :1, :, :].repeat((1, 1, padding_length, 1, 1)) pieces = [ cached, x ] if is_end and not causal: @@ -83,7 +89,7 @@ class CausalConv3d(nn.Module): elif is_end: self.temporal_cache_state[tid] = (None, True) - return self.conv(x) if x.shape[2] >= self.time_kernel_size else x[:, :, :0, :, :] + return self.conv(x) if x.shape[2] >= self.time_kernel_size else self._empty_output(x) @property def weight(self): diff --git a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py index 5975015e2..5d0eec5b8 100644 --- a/comfy/ldm/lightricks/vae/causal_video_autoencoder.py +++ b/comfy/ldm/lightricks/vae/causal_video_autoencoder.py @@ -390,10 +390,10 @@ class Decoder(nn.Module): # Compute output channel to be product of all channel-multiplier blocks output_channel = base_channels - for block_name, block_params in list(reversed(blocks)): + for block_name, block_params in blocks: block_params = block_params if isinstance(block_params, dict) else {} if block_name == "res_x_y": - output_channel = output_channel * block_params.get("multiplier", 2) + output_channel = block_params.get("in_channels", output_channel * block_params.get("multiplier", 2)) if block_name == "compress_all": output_channel = output_channel * block_params.get("multiplier", 1) if block_name == "compress_space": @@ -432,7 +432,7 @@ class Decoder(nn.Module): spatial_padding_mode=spatial_padding_mode, ) elif block_name == "res_x_y": - output_channel = output_channel // block_params.get("multiplier", 2) + output_channel = block_params.get("out_channels", output_channel // block_params.get("multiplier", 2)) block = ResnetBlock3D( dims=dims, in_channels=input_channel, diff --git a/comfy/ldm/lumina/model.py b/comfy/ldm/lumina/model.py index d0ee97d33..cdf03b2b5 100644 --- a/comfy/ldm/lumina/model.py +++ b/comfy/ldm/lumina/model.py @@ -6,6 +6,9 @@ import torch import torch.nn as nn import torch.nn.functional as F import comfy.ldm.common_dit +import comfy.model_management +import comfy.ops +import comfy.quant_ops from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder from comfy.ldm.modules.attention import optimized_attention_masked @@ -97,6 +100,7 @@ class JointAttention(nn.Module): self.n_local_kv_heads = self.n_kv_heads self.n_rep = self.n_local_heads // self.n_local_kv_heads self.head_dim = dim // n_heads + self.qk_norm = qk_norm self.qkv = operation_settings.get("operations").Linear( dim, @@ -151,10 +155,21 @@ class JointAttention(nn.Module): xk = xk.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) xv = xv.view(bsz, seqlen, self.n_local_kv_heads, self.head_dim) - xq = self.q_norm(xq) - xk = self.k_norm(xk) - - xq, xk = apply_rope(xq, xk, freqs_cis) + if self.qk_norm and not comfy.model_management.in_training: + q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, xq, offloadable=True) + k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, xk, offloadable=True) + epsilon = self.q_norm.eps if self.q_norm.eps is not None else torch.finfo(torch.float32).eps + if self.n_local_heads == self.n_local_kv_heads: + xq, xk = comfy.quant_ops.ck.rms_rope(xq, xk, freqs_cis, q_scale, k_scale, epsilon) + else: + xq = comfy.quant_ops.ck.rms_rope1(xq, freqs_cis, q_scale, epsilon) + xk = comfy.quant_ops.ck.rms_rope1(xk, freqs_cis, k_scale, epsilon) + comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream) + comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream) + else: + xq = self.q_norm(xq) + xk = self.k_norm(xk) + xq, xk = apply_rope(xq, xk, freqs_cis) n_rep = self.n_local_heads // self.n_local_kv_heads if n_rep >= 1: diff --git a/comfy/ldm/mage_flow/model.py b/comfy/ldm/mage_flow/model.py new file mode 100644 index 000000000..ac29bb610 --- /dev/null +++ b/comfy/ldm/mage_flow/model.py @@ -0,0 +1,186 @@ +# Mage-Flow (https://github.com/microsoft/Mage) native-resolution MMDiT (MIT) +# Architecture is a 12-layer variant of the Qwen-Image double-stream block with +# patch_size=1 (no 2x2 packing), unrotated text tokens and a bf16-rounded +# timestep frequency table. +import math +import torch +import torch.nn as nn +from typing import Optional, Tuple + +from comfy.ldm.lightricks.model import TimestepEmbedding +from comfy.ldm.flux.layers import EmbedND +from comfy.ldm.qwen_image.model import QwenImageTransformerBlock, LastLayer +import comfy.patcher_extension + + +class MageTimestepProjEmbeddings(nn.Module): + def __init__(self, embedding_dim, dtype=None, device=None, operations=None): + super().__init__() + self.timestep_embedder = TimestepEmbedding( + in_channels=256, time_embed_dim=embedding_dim, + dtype=dtype, device=device, operations=operations + ) + + def forward(self, timestep, hidden_states): + half_dim = 128 + exponent = -math.log(10000) * torch.arange(half_dim, dtype=torch.float32, device=timestep.device) / half_dim + emb = torch.exp(exponent).to(timestep.dtype) + emb = timestep[:, None].float() * emb[None, :] + emb = 1000.0 * emb + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) # flip_sin_to_cos + return self.timestep_embedder(emb.to(dtype=hidden_states.dtype)) + + +class MageFlowTransformer2DModel(nn.Module): + def __init__( + self, + in_channels: int = 128, + out_channels: Optional[int] = 128, + num_layers: int = 12, + attention_head_dim: int = 128, + num_attention_heads: int = 24, + joint_attention_dim: int = 2560, + axes_dims_rope: Tuple[int, int, int] = (16, 56, 56), + image_model=None, + dtype=None, + device=None, + operations=None, + ): + super().__init__() + self.dtype = dtype + self.patch_size = 1 + self.in_channels = in_channels + self.out_channels = out_channels or in_channels + self.inner_dim = num_attention_heads * attention_head_dim + + self.pe_embedder = EmbedND(dim=attention_head_dim, theta=10000, axes_dim=list(axes_dims_rope)) + + self.time_text_embed = MageTimestepProjEmbeddings(embedding_dim=self.inner_dim, dtype=dtype, device=device, operations=operations) + + self.txt_norm = operations.RMSNorm(joint_attention_dim, eps=1e-6, dtype=dtype, device=device) + self.img_in = operations.Linear(in_channels, self.inner_dim, dtype=dtype, device=device) + self.txt_in = operations.Linear(joint_attention_dim, self.inner_dim, dtype=dtype, device=device) + + self.transformer_blocks = nn.ModuleList([ + QwenImageTransformerBlock( + dim=self.inner_dim, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + dtype=dtype, + device=device, + operations=operations + ) + for _ in range(num_layers) + ]) + + self.norm_out = LastLayer(self.inner_dim, self.inner_dim, dtype=dtype, device=device, operations=operations) + self.proj_out = operations.Linear(self.inner_dim, self.out_channels, bias=True, dtype=dtype, device=device) + + def process_img(self, x, index=0): + # patch_size=1: tokens are raw latent pixels, no 2x2 packing. + bs, c, h, w = x.shape + hidden_states = x.movedim(1, -1).reshape(bs, h * w, c) + + img_ids = torch.zeros((h, w, 3), device=x.device) + # Frame axis: positive image index (0 = target, 1..N = reference images). + img_ids[:, :, 0] = index + # Mage scale_rope centering: positions [-ceil(n/2), floor(n/2)), i.e. + # offset by (n - n//2). Differs from Qwen-Image's -(n//2) for odd sizes. + img_ids[:, :, 1] = img_ids[:, :, 1] + torch.arange(h, device=x.device)[:, None] - (h - h // 2) + img_ids[:, :, 2] = img_ids[:, :, 2] + torch.arange(w, device=x.device)[None, :] - (w - w // 2) + return hidden_states, img_ids.reshape(h * w, 3).unsqueeze(0).expand(bs, -1, -1), (h, w) + + def forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(x, timestep, context, attention_mask, ref_latents, transformer_options, **kwargs) + + def _forward(self, x, timestep, context, attention_mask=None, ref_latents=None, transformer_options={}, control=None, **kwargs): + if attention_mask is not None and not torch.is_floating_point(attention_mask): + attention_mask = (attention_mask - 1).to(x.dtype) * torch.finfo(x.dtype).max + + hidden_states, img_ids, orig_shape = self.process_img(x) + num_embeds = hidden_states.shape[1] + + if ref_latents is not None: + ref_num_tokens = [] + index = 0 + for ref in ref_latents: + index += 1 + kontext, kontext_ids, _ = self.process_img(ref, index=index) + hidden_states = torch.cat([hidden_states, kontext], dim=1) + img_ids = torch.cat([img_ids, kontext_ids], dim=1) + ref_num_tokens.append(kontext.shape[1]) + transformer_options = transformer_options.copy() + transformer_options["reference_image_num_tokens"] = ref_num_tokens + + # Text tokens are not rotated in Mage-Flow: RoPE at position 0 is the + # identity rotation. + txt_ids = torch.zeros((x.shape[0], context.shape[1], 3), device=x.device) + + hidden_states = self.img_in(hidden_states) + context = self.txt_norm(context) + context = self.txt_in(context) + + temb = self.time_text_embed(timestep, hidden_states) + + patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) + blocks_replace = patches_replace.get("dit", {}) + + if "post_input" in patches: + for p in patches["post_input"]: + out = p({"img": hidden_states, "txt": context, "img_ids": img_ids, "txt_ids": txt_ids, "transformer_options": transformer_options}) + hidden_states = out["img"] + context = out["txt"] + img_ids = out["img_ids"] + txt_ids = out["txt_ids"] + + ids = torch.cat((txt_ids, img_ids), dim=1) + image_rotary_emb = self.pe_embedder(ids).contiguous() + del ids, txt_ids, img_ids + + transformer_options["total_blocks"] = len(self.transformer_blocks) + transformer_options["block_type"] = "double" + for i, block in enumerate(self.transformer_blocks): + transformer_options["block_index"] = i + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["txt"], out["img"] = block(hidden_states=args["img"], encoder_hidden_states=args["txt"], encoder_hidden_states_mask=attention_mask, temb=args["vec"], image_rotary_emb=args["pe"], transformer_options=args["transformer_options"]) + return out + out = blocks_replace[("double_block", i)]({"img": hidden_states, "txt": context, "vec": temb, "pe": image_rotary_emb, "transformer_options": transformer_options}, {"original_block": block_wrap}) + hidden_states = out["img"] + context = out["txt"] + else: + context, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=context, + encoder_hidden_states_mask=attention_mask, + temb=temb, + image_rotary_emb=image_rotary_emb, + transformer_options=transformer_options, + ) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": hidden_states, "txt": context, "x": x, "block_index": i, "transformer_options": transformer_options}) + hidden_states = out["img"] + context = out["txt"] + + if control is not None: # Controlnet + control_i = control.get("input") + if i < len(control_i): + add = control_i[i] + if add is not None: + hidden_states[:, :add.shape[1]] += add + + hidden_states = self.norm_out(hidden_states, temb) + hidden_states = self.proj_out(hidden_states) + + hidden_states = hidden_states[:, :num_embeds] + h, w = orig_shape + return hidden_states.reshape(x.shape[0], h, w, self.out_channels).movedim(-1, 1) diff --git a/comfy/ldm/mage_flow/vae.py b/comfy/ldm/mage_flow/vae.py new file mode 100644 index 000000000..e6e21b99f --- /dev/null +++ b/comfy/ldm/mage_flow/vae.py @@ -0,0 +1,477 @@ +# Mage-VAE (https://github.com/microsoft/Mage) (MIT) +# Symmetric one-step diffusion codec: DConvEncoder (image -> 128ch latent) and +# DConvDenoiser + CoD Decoder (latent -> image). 16x downsample, latents in the +# Flux.2-VAE-anchored space (no patch packing, no BN normalization). +# Both encode and decode are single forward passes at t=0. +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops +from comfy.ldm.modules.diffusionmodules.model import vae_attention + +ops = comfy.ops.disable_weight_init + + +def nonlinearity(x): + return torch.nn.functional.silu(x) + + +def Normalize(in_channels): + return ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + + +def modulate(x, shift, scale): + if x.dim() == 4: + b, c = x.shape[:2] + return x * (1 + scale.view(b, c, 1, 1)) + shift.view(b, c, 1, 1) + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + +class LayerNorm2d(ops.LayerNorm): + def __init__(self, num_channels, eps=1e-6, affine=True): + super().__init__(num_channels, eps=eps, elementwise_affine=affine) + + def forward(self, x): + x = x.permute(0, 2, 3, 1).contiguous() + x = super().forward(x) + return x.permute(0, 3, 1, 2).contiguous() + + +class TimestepEmbedder(nn.Module): + """DConv-style timestep MLP (max_period=10000, freq_size=256).""" + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__() + self.mlp = nn.Sequential( + ops.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + ops.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + @staticmethod + def timestep_embedding(t, dim, max_period=10000): + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half + ).to(t.device) + args = t[:, None].float() * freqs[None] + emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1) + return emb + + def forward(self, t, dtype): + emb = self.timestep_embedding(t, self.frequency_embedding_size) + return self.mlp(emb.to(dtype)) + + +class BottleneckPatchEmbed(nn.Module): + """Image patch embed concatenated with a per-patch conditioning vector.""" + + def __init__(self, patch_size=16, in_chans=3, pca_dim=128, embed_dim=384, bias=True): + super().__init__() + self.proj1 = ops.Conv2d(in_chans, pca_dim, kernel_size=patch_size, stride=patch_size, bias=False) + self.proj2 = ops.Conv2d(pca_dim + embed_dim, embed_dim, kernel_size=1, bias=bias) + + def forward(self, x, cond): + return self.proj2(torch.cat([self.proj1(x), cond], dim=1)) + + +class DiCoBlock(nn.Module): + """DConv block with adaLN modulation.""" + + def __init__(self, hidden_size, mlp_ratio=4.0): + super().__init__() + self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True) + self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + + self.ca = nn.Sequential( + nn.AdaptiveAvgPool2d(1), + ops.Conv2d(hidden_size, hidden_size, 1, bias=True), + nn.Sigmoid(), + ) + + ffn = int(mlp_ratio * hidden_size) + self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True) + self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True) + + self.norm1 = LayerNorm2d(hidden_size, affine=False) + self.norm2 = LayerNorm2d(hidden_size, affine=False) + + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + ops.Linear(hidden_size, 6 * hidden_size, bias=True), + ) + + def forward(self, inp, c): + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1) + x = modulate(self.norm1(inp), shift_msa, scale_msa) + x = F.gelu(self.conv2(self.conv1(x))) + x = x * self.ca(x) + x = self.conv3(x) + x = inp + gate_msa[..., None, None] * x + x = x + gate_mlp[..., None, None] * self.conv5( + F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp))) + ) + return x + + +class EncoderDiCoBlock(nn.Module): + """DiCoBlock without adaLN, for the encoder head pathway.""" + + def __init__(self, hidden_size, mlp_ratio=4.0): + super().__init__() + self.conv1 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.conv2 = ops.Conv2d(hidden_size, hidden_size, 3, padding=1, groups=hidden_size, bias=True) + self.conv3 = ops.Conv2d(hidden_size, hidden_size, 1, bias=True) + self.ca = nn.Sequential( + nn.AdaptiveAvgPool2d(1), + ops.Conv2d(hidden_size, hidden_size, 1, bias=True), + nn.Sigmoid(), + ) + ffn = int(mlp_ratio * hidden_size) + self.conv4 = ops.Conv2d(hidden_size, ffn, 1, bias=True) + self.conv5 = ops.Conv2d(ffn, hidden_size, 1, bias=True) + self.norm1 = LayerNorm2d(hidden_size) + self.norm2 = LayerNorm2d(hidden_size) + + def forward(self, inp): + x = self.norm1(inp) + x = F.gelu(self.conv2(self.conv1(x))) + x = x * self.ca(x) + x = self.conv3(x) + x = inp + x + return x + self.conv5(F.gelu(self.conv4(self.norm2(x)))) + + +class NerfEmbedder(nn.Module): + """Patch-position embedder used by the DConv decoder x-pathway.""" + + def __init__(self, in_channels, hidden_size_input, max_freqs=8): + super().__init__() + self.max_freqs = max_freqs + self.embedder = nn.Sequential( + ops.Linear(in_channels + max_freqs ** 2, hidden_size_input, bias=True), + ) + + def fetch_pos(self, patch_size, device, dtype): + pos = torch.linspace(0, 1, patch_size, device=device, dtype=dtype) + pos_y, pos_x = torch.meshgrid(pos, pos, indexing="ij") + pos_x = pos_x.reshape(-1, 1, 1) + pos_y = pos_y.reshape(-1, 1, 1) + freqs = torch.linspace(0, self.max_freqs, self.max_freqs, dtype=dtype, device=device) + fx = freqs[None, :, None] + fy = freqs[None, None, :] + coeffs = (1 + fx * fy) ** -1 + dct_x = torch.cos(pos_x * fx * torch.pi) + dct_y = torch.cos(pos_y * fy * torch.pi) + return (dct_x * dct_y * coeffs).view(1, -1, self.max_freqs ** 2) + + def forward(self, x): + B, P2, _ = x.shape + ps = int(P2 ** 0.5) + dct = self.fetch_pos(ps, x.device, x.dtype).expand(B, -1, -1) + return self.embedder(torch.cat([x, dct], dim=-1)) + + +class NerfFinalLayer(nn.Module): + def __init__(self, hidden_size, out_channels): + super().__init__() + self.norm = ops.RMSNorm(hidden_size, eps=1e-6) + self.linear = ops.Linear(hidden_size, out_channels, bias=True) + + def forward(self, x): + return self.linear(self.norm(x)) + + +class MLPResBlock(nn.Module): + def __init__(self, channels): + super().__init__() + self.in_ln = ops.LayerNorm(channels, eps=1e-6) + self.mlp = nn.Sequential( + ops.Linear(channels, channels, bias=True), + nn.SiLU(), + ops.Linear(channels, channels, bias=True), + ) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + ops.Linear(channels, 3 * channels, bias=True), + ) + + def forward(self, x, y): + shift, scale, gate = self.adaLN_modulation(y).chunk(3, dim=-1) + h = self.in_ln(x) * (1 + scale) + shift + return x + gate * self.mlp(h) + + +class SimpleMLPAdaLN(nn.Module): + """Final small MLP that maps NerfEmbedder features to per-patch RGB.""" + + def __init__(self, in_channels, model_channels, out_channels, z_channels, num_res_blocks, patch_size): + super().__init__() + self.in_channels = in_channels + self.model_channels = model_channels + self.out_channels = out_channels + self.num_res_blocks = num_res_blocks + self.patch_size = patch_size + + self.cond_embed = ops.Linear(z_channels, patch_size ** 2 * model_channels) + self.input_proj = ops.Linear(in_channels, model_channels) + + self.res_blocks = nn.ModuleList(MLPResBlock(model_channels) for _ in range(num_res_blocks)) + + def forward(self, x, c): + x = self.input_proj(x) + c = self.cond_embed(c).reshape(c.shape[0], self.patch_size ** 2, -1) + for block in self.res_blocks: + x = block(x, c) + return x + + +class ResnetBlock(nn.Module): + """GroupNorm + Conv ResBlock used by the CoD Decoder.""" + + def __init__(self, *, in_channels, out_channels=None): + super().__init__() + out_channels = out_channels or in_channels + self.in_channels = in_channels + self.out_channels = out_channels + + self.norm1 = Normalize(in_channels) + self.conv1 = ops.Conv2d(in_channels, out_channels, 3, padding=1) + self.norm2 = Normalize(out_channels) + self.conv2 = ops.Conv2d(out_channels, out_channels, 3, padding=1) + if in_channels != out_channels: + self.nin_shortcut = ops.Conv2d(in_channels, out_channels, 1) + + def forward(self, x): + h = self.conv1(nonlinearity(self.norm1(x))) + h = self.conv2(nonlinearity(self.norm2(h))) + if self.in_channels != self.out_channels: + x = self.nin_shortcut(x) + return x + h + + +class AttnBlock(nn.Module): + """Patched (windowed) self-attention used by the CoD Decoder.""" + + def __init__(self, in_channels, patch_size=32): + super().__init__() + self.in_channels = in_channels + self.patch_size = patch_size + self.norm = Normalize(in_channels) + self.q = ops.Conv2d(in_channels, in_channels, 1) + self.k = ops.Conv2d(in_channels, in_channels, 1) + self.v = ops.Conv2d(in_channels, in_channels, 1) + self.proj_out = ops.Conv2d(in_channels, in_channels, 1) + # VAE attention selection: full-precision backends only (no sage/quantized attention) + self.optimized_attention = vae_attention() + + def forward(self, x): + h_ = self.norm(x) + Q = self.q(h_) + K = self.k(h_) + V = self.v(h_) + + d = self.patch_size + b, c, H, W = Q.shape + pad_h = (d - H % d) % d + pad_w = (d - W % d) % d + if pad_h or pad_w: + Q = F.pad(Q, (0, pad_w, 0, pad_h), mode="replicate") + K = F.pad(K, (0, pad_w, 0, pad_h), mode="replicate") + V = F.pad(V, (0, pad_w, 0, pad_h), mode="replicate") + _, _, H_pad, W_pad = Q.shape + nph, npw = H_pad // d, W_pad // d + np_ = nph * npw + + def to_patches(t): + return (t.reshape(b, c, nph, d, npw, d) + .permute(0, 2, 4, 1, 3, 5) + .reshape(b * np_, c, d * d)) + + # [b*np, c, d*d]: attention over the d*d spatial positions of each window + Q = to_patches(Q) + K = to_patches(K) + V = to_patches(V) + + h_ = self.optimized_attention(Q, K, V) + h_ = h_.reshape(b, nph, npw, c, d, d).permute(0, 3, 1, 4, 2, 5).reshape(b, c, H_pad, W_pad) + if pad_h or pad_w: + h_ = h_[:, :, :H, :W] + return x + self.proj_out(h_) + + +class CoDDecoder(nn.Module): + """CoD Decoder: latent -> conditioning features for the denoiser (ds=16, light).""" + + def __init__(self, out_ch=384, z_ch=128): + super().__init__() + self.conv_in = ops.Conv2d(z_ch, out_ch, kernel_size=3, stride=1, padding=1) + self.block = nn.Sequential( + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + AttnBlock(out_ch, patch_size=32), + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + AttnBlock(out_ch, patch_size=32), + ResnetBlock(in_channels=out_ch, out_channels=out_ch), + ) + self.norm_out = Normalize(out_ch) + self.conv_out = ops.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1) + self.ada = nn.Identity() + + def forward(self, z): + h = self.block(self.conv_in(z)) + h = self.conv_out(nonlinearity(self.norm_out(h))) + return self.ada(h) + + +class DConvEncoder(nn.Module): + """DConvEncoder: image -> packed (mean, logvar) latent.""" + + def __init__( + self, + z_ch=128, + hidden_size=384, + num_blocks=21, + patch_size=16, + mlp_ratio=4.0, + head_size=768, + num_head_blocks=2, + out_ch_mult=2, + ): + super().__init__() + self.z_ch = z_ch + self.patch_size = patch_size + self.patch_cond_embed = ops.Conv2d(3, head_size, kernel_size=patch_size, stride=patch_size, bias=True) + self.head_blocks = nn.ModuleList([ + EncoderDiCoBlock(head_size, mlp_ratio=mlp_ratio) for _ in range(num_head_blocks) + ]) + self.proj_down = ops.Conv2d(head_size, hidden_size, kernel_size=1, bias=True) + self.z_proj = ops.Conv2d(z_ch, hidden_size, kernel_size=1, bias=True) + self.fuse_proj = ops.Conv2d(hidden_size * 2, hidden_size, kernel_size=1, bias=True) + self.t_embedder = TimestepEmbedder(hidden_size) + self.blocks = nn.ModuleList([ + DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_blocks) + ]) + self.norm_out = LayerNorm2d(hidden_size) + self.proj_out = ops.Conv2d(hidden_size, z_ch * out_ch_mult, kernel_size=1, bias=True) + + def forward_pred(self, z_t, t, y): + cond = self.patch_cond_embed(y) + for block in self.head_blocks: + cond = block(cond) + cond = self.proj_down(cond) + + s = self.fuse_proj(torch.cat([cond, self.z_proj(z_t)], dim=1)) + c = self.t_embedder(t.view(-1), y.dtype) + for block in self.blocks: + s = block(s, c) + return self.proj_out(self.norm_out(s)) + + +class YEmbedder(nn.Module): + """Holds only the CoD decoder (the original Flux2-VAE encoder side is dropped at load).""" + + def __init__(self, ch=384, z_ch=128): + super().__init__() + self.decoder = CoDDecoder(out_ch=ch, z_ch=z_ch) + + +class DConvDenoiser(nn.Module): + """One-step DConv denoiser: latent (via cond) + zero noise -> reconstructed image.""" + + def __init__( + self, + patch_size=16, + in_channels=3, + hidden_size=384, + hidden_size_x=32, + mlp_ratio=4.0, + num_blocks=24, + num_cond_blocks=21, + bottleneck_dim=128, + ): + super().__init__() + self.in_channels = in_channels + self.patch_size = patch_size + self.hidden_size = hidden_size + self.num_cond_blocks = num_cond_blocks + + self.t_embedder = TimestepEmbedder(hidden_size) + self.y_embedder_x = ops.Conv2d(hidden_size, hidden_size_x * patch_size ** 2, 1, 1, 0) + self.x_embedder = NerfEmbedder(in_channels + hidden_size_x, hidden_size_x, max_freqs=8) + self.s_embedder = BottleneckPatchEmbed(patch_size, in_channels, bottleneck_dim, hidden_size, bias=True) + self.blocks = nn.ModuleList([ + DiCoBlock(hidden_size, mlp_ratio=mlp_ratio) for _ in range(num_cond_blocks) + ]) + self.dec_net = SimpleMLPAdaLN( + in_channels=hidden_size_x, + model_channels=hidden_size_x, + out_channels=in_channels, + z_channels=hidden_size, + num_res_blocks=num_blocks - num_cond_blocks, + patch_size=patch_size, + ) + self.final_layer = NerfFinalLayer(hidden_size_x, in_channels) + self.y_embedder = YEmbedder(ch=hidden_size, z_ch=bottleneck_dim) + + def forward(self, x, t, cond): + b, _, h, w = x.shape + c = self.t_embedder(t.view(-1), x.dtype) + + s = self.s_embedder(x, cond) + for block in self.blocks: + s = block(s, c) + + length = s.shape[-2] * s.shape[-1] + s = s.permute(0, 2, 3, 1).reshape(-1, self.hidden_size) + + x = torch.nn.functional.unfold(x, kernel_size=self.patch_size, stride=self.patch_size) + x = torch.cat([x, self.y_embedder_x(cond).flatten(2)], dim=1) + x = x.reshape(b, -1, self.patch_size ** 2, length).permute(0, 3, 2, 1).flatten(0, 1) + x = self.x_embedder(x) + + x = self.dec_net(x, s) + x = self.final_layer(x) + x = x.transpose(1, 2).reshape(b, length, -1) + return torch.nn.functional.fold( + x.transpose(1, 2).contiguous(), (h, w), + kernel_size=self.patch_size, stride=self.patch_size, + ) + + +class MageVAE(nn.Module): + """ + Encode: DConvEncoder (one-step at t=0) -> posterior mean [B, 128, H/16, W/16] + Decode: DConvDenoiser + CoD Decoder -> image [B, 3, H, W] in [-1, 1] + """ + + latent_channels = 128 + downsample_factor = 16 + + def __init__(self): + super().__init__() + self.dconv_encoder = DConvEncoder() + self.decoder_model = DConvDenoiser() + + def encode(self, x): + B, _, H, W = x.shape + ps = self.dconv_encoder.patch_size + z_t = torch.zeros(B, self.dconv_encoder.z_ch, H // ps, W // ps, device=x.device, dtype=x.dtype) + t = torch.zeros(B, device=x.device, dtype=x.dtype) + out = self.dconv_encoder.forward_pred(z_t, t, x) + return out[:, : self.latent_channels] # posterior mean (sample_posterior=False) + + def decode(self, z): + cond = self.decoder_model.y_embedder.decoder(z) + B = z.shape[0] + H = z.shape[2] * self.downsample_factor + W = z.shape[3] * self.downsample_factor + noise = torch.zeros(B, 3, H, W, device=z.device, dtype=z.dtype) + t = torch.zeros(B, device=z.device, dtype=z.dtype) + return self.decoder_model.forward(noise, t, cond) diff --git a/comfy/ldm/minimax/audio_vae.py b/comfy/ldm/minimax/audio_vae.py new file mode 100644 index 000000000..033ae6966 --- /dev/null +++ b/comfy/ldm/minimax/audio_vae.py @@ -0,0 +1,442 @@ +# MiniMax H3 audio VAE: DAC-lineage waveform encoder + BigVGAN decoder. +# Weight-norm parametrizations are folded into plain conv weights, so this +# module uses ordinary ops.Conv1d / ops.ConvTranspose1d and loads the converted +# checkpoint (plain "*.weight" tensors) with strict=True. +# +# Lineage / licenses of the reference implementation: +# DAC encoder: descript-audio-codec (MIT) +# BigVGAN decoder: NVIDIA BigVGAN (MIT), adapted from hifi-gan (MIT) +# Alias-free ops: junjun3518/alias-free-torch (Apache-2.0), julius (MIT) + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops + +ops = comfy.ops.disable_weight_init + + +# Snake activations + +def snake(x, alpha, beta): + # x + 1/beta * sin^2(alpha * x) + t = torch.sin(alpha * x) + return t.mul_(t).mul_((beta + 1e-9).reciprocal()).add_(x) + + +class Snake1d(nn.Module): + """Snake activation with per-channel alpha (encoder side).""" + + def __init__(self, channels): + super().__init__() + self.alpha = nn.Parameter(torch.empty(1, channels, 1)) + + def forward(self, x): + return snake(x, self.alpha, self.alpha) + + +class SnakeBeta(nn.Module): + """SnakeBeta := x + 1/beta * sin^2(alpha * x); alpha/beta stored in log scale.""" + + def __init__(self, in_features): + super().__init__() + self.alpha = nn.Parameter(torch.empty(in_features)) + self.beta = nn.Parameter(torch.empty(in_features)) + + def forward(self, x): + alpha = torch.exp(self.alpha).view(1, -1, 1) + beta = torch.exp(self.beta).view(1, -1, 1) + return snake(x, alpha, beta) + + +# Alias-free (anti-aliased) activation: kaiser-windowed sinc resampling + +def kaiser_sinc_filter1d(cutoff, half_width, kernel_size): + # returns filter [1, 1, kernel_size] + even = kernel_size % 2 == 0 + half_size = kernel_size // 2 + + # kaiser window design + delta_f = 4 * half_width + A = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95 + if A > 50.0: + beta = 0.1102 * (A - 8.7) + elif A >= 21.0: + beta = 0.5842 * (A - 21) ** 0.4 + 0.07886 * (A - 21.0) + else: + beta = 0.0 + window = torch.kaiser_window(kernel_size, beta=beta, periodic=False) + + if even: + time = torch.arange(-half_size, half_size) + 0.5 + else: + time = torch.arange(kernel_size) - half_size + + filter_ = 2 * cutoff * window * torch.sinc(2 * cutoff * time) + # Normalize filter to have sum = 1, otherwise there is a small leakage of + # the constant component in the input signal. + filter_ /= filter_.sum() + return filter_.view(1, 1, kernel_size) + + +class UpSample1d(nn.Module): + def __init__(self, ratio=2, kernel_size=12): + super().__init__() + self.ratio = ratio + self.stride = ratio + self.pad = kernel_size // ratio - 1 + self.pad_left = self.pad * ratio + (kernel_size - ratio) // 2 + self.pad_right = self.pad * ratio + (kernel_size - ratio + 1) // 2 + self.register_buffer( + "filter", + kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=kernel_size), + ) + + def forward(self, x): + _, C, _ = x.shape + x = F.pad(x, (self.pad, self.pad), mode="replicate") + x = F.conv_transpose1d(x, self.filter.expand(C, -1, -1).to(x.dtype), stride=self.stride, groups=C).mul_(self.ratio) + x = x[..., self.pad_left:-self.pad_right] + return x + + +class LowPassFilter1d(nn.Module): + def __init__(self, cutoff=0.5, half_width=0.6, stride=1, kernel_size=12): + super().__init__() + self.pad_left = kernel_size // 2 - int(kernel_size % 2 == 0) + self.pad_right = kernel_size // 2 + self.stride = stride + self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size)) + + def forward(self, x): + _, C, _ = x.shape + x = F.pad(x, (self.pad_left, self.pad_right), mode="replicate") + return F.conv1d(x, self.filter.expand(C, -1, -1).to(x.dtype), stride=self.stride, groups=C) + + +class DownSample1d(nn.Module): + def __init__(self, ratio=2, kernel_size=12): + super().__init__() + self.ratio = ratio + self.kernel_size = kernel_size + self.lowpass = LowPassFilter1d( + cutoff=0.5 / ratio, + half_width=0.6 / ratio, + stride=ratio, + kernel_size=self.kernel_size, + ) + + def forward(self, x): + return self.lowpass(x) + + +class Activation1d(nn.Module): + """upsample x2 -> pointwise activation -> downsample x2 (anti-aliased).""" + + def __init__(self, activation, up_ratio=2, down_ratio=2, up_kernel_size=12, down_kernel_size=12): + super().__init__() + self.act = activation + self.upsample = UpSample1d(up_ratio, up_kernel_size) + self.downsample = DownSample1d(down_ratio, down_kernel_size) + + def forward(self, x): + x = self.upsample(x) + x = self.act(x) + x = self.downsample(x) + return x + + +# DAC encoder + +class ResidualUnit(nn.Module): + def __init__(self, dim=16, dilation=1): + super().__init__() + pad = ((7 - 1) * dilation) // 2 + self.block = nn.Sequential( + Snake1d(dim), + ops.Conv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad), + Snake1d(dim), + ops.Conv1d(dim, dim, kernel_size=1), + ) + + def forward(self, x): + y = self.block(x) + pad = (x.shape[-1] - y.shape[-1]) // 2 + if pad > 0: + x = x[..., pad:-pad] + return y.add_(x) + + +class EncoderBlock(nn.Module): + def __init__(self, dim=16, stride=1): + super().__init__() + self.block = nn.Sequential( + ResidualUnit(dim // 2, dilation=1), + ResidualUnit(dim // 2, dilation=3), + ResidualUnit(dim // 2, dilation=9), + Snake1d(dim // 2), + ops.Conv1d( + dim // 2, + dim, + kernel_size=2 * stride, + stride=stride, + padding=math.ceil(stride / 2), + ), + ) + + def forward(self, x): + return self.block(x) + + +class Encoder(nn.Module): + def __init__(self, d_model=64, strides=(2, 4, 4, 5, 5), d_latent=2048): + super().__init__() + block = [ops.Conv1d(1, d_model, kernel_size=7, padding=3)] + for stride in strides: + d_model *= 2 + block += [EncoderBlock(d_model, stride=stride)] + block += [ + Snake1d(d_model), + ops.Conv1d(d_model, d_latent, kernel_size=3, padding=1), + ] + self.block = nn.Sequential(*block) + + def forward(self, x): + return self.block(x) + + +# Attention projection (encoder posterior head) + +class GeGluMlp(nn.Module): + def __init__(self, in_features, hidden_features): + super().__init__() + self.norm = ops.LayerNorm(in_features) + self.act = nn.GELU(approximate="tanh") + self.w0 = ops.Linear(in_features, hidden_features) + self.w1 = ops.Linear(in_features, hidden_features) + self.w2 = ops.Linear(hidden_features, in_features) + + def forward(self, x): + x = self.norm(x) + return self.w2(self.act(self.w0(x)).mul_(self.w1(x))) + + +class CausalAttention(nn.Module): + def __init__(self, in_dim, out_dim, num_heads): + super().__init__() + self.head_dim = in_dim // num_heads + self.num_heads = num_heads + self.out_dim = out_dim + self.qkv = ops.Linear(in_dim, in_dim * 3, bias=False) + self.q_bias = nn.Parameter(torch.empty(in_dim)) + self.v_bias = nn.Parameter(torch.empty(in_dim)) + self.register_buffer("zero_k_bias", torch.empty(in_dim)) + self.proj = ops.Linear(out_dim, out_dim) + + def forward(self, x): + B, N, C = x.shape + weight, _, offload_stream = comfy.ops.cast_bias_weight(self.qkv, x, offloadable=True) + qkv = F.linear(x, weight=weight, bias=torch.cat((self.q_bias, self.zero_k_bias, self.v_bias))) + comfy.ops.uncast_bias_weight(self.qkv, weight, None, offload_stream) + q, k, v = qkv.reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4).unbind(0) + + # mean over heads then pool down to the latent width (in_dim >> out_dim) + x = comfy.ops.scaled_dot_product_attention(q, k, v, is_causal=True) + x = F.adaptive_avg_pool1d(torch.mean(x, dim=1), self.out_dim) + return self.proj(x) + + +class AttnProjection(nn.Module): + def __init__(self, in_dim, out_dim, num_heads, mlp_ratio=2): + super().__init__() + self.norm1 = ops.LayerNorm(in_dim) + self.attn = CausalAttention(in_dim, out_dim, num_heads) + self.proj = ops.Linear(in_dim, out_dim) + self.norm3 = ops.LayerNorm(in_dim) + + self.norm2 = ops.LayerNorm(out_dim) + hidden_dim = int(out_dim * mlp_ratio) + self.mlp = GeGluMlp(in_features=out_dim, hidden_features=hidden_dim) + + def forward(self, x): + # x: [B, T, in_dim] + x = self.proj(self.norm3(x)).add_(self.attn(self.norm1(x))) + return x.add_(self.mlp(self.norm2(x))) + + +# BigVGAN decoder + +def get_padding(kernel_size, dilation=1): + return int((kernel_size * dilation - dilation) / 2) + + +class AMPBlock1(nn.Module): + def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)): + super().__init__() + self.convs1 = nn.ModuleList( + [ + ops.Conv1d(channels, channels, kernel_size, stride=1, dilation=d, padding=get_padding(kernel_size, d)) + for d in dilation + ] + ) + self.convs2 = nn.ModuleList( + [ + ops.Conv1d(channels, channels, kernel_size, stride=1, dilation=1, padding=get_padding(kernel_size, 1)) + for _ in range(len(dilation)) + ] + ) + self.num_layers = len(self.convs1) + len(self.convs2) + self.activations = nn.ModuleList( + [Activation1d(activation=SnakeBeta(channels)) for _ in range(self.num_layers)] + ) + + def forward(self, x): + acts1, acts2 = self.activations[::2], self.activations[1::2] + for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2): + xt = a1(x) + xt = c1(xt) + xt = a2(xt) + xt = c2(xt) + x = xt.add_(x) + return x + + +class BigVGAN(nn.Module): + """BigVGAN vocoder (MiniMax H3 32 kHz configuration). + + use_bias_at_final=False, use_tanh_at_final=False (output clamped to [-1, 1]). + """ + + def __init__( + self, + num_mels=2048, + upsample_initial_channel=1024, + upsample_rates=(5, 5, 2, 2, 2, 2, 2), + upsample_kernel_sizes=(9, 9, 4, 4, 4, 4, 4), + resblock_kernel_sizes=(3, 7, 11), + resblock_dilation_sizes=((1, 3, 5), (1, 3, 5), (1, 3, 5)), + ): + super().__init__() + self.num_kernels = len(resblock_kernel_sizes) + self.num_upsamples = len(upsample_rates) + + self.conv_pre = ops.Conv1d(num_mels, upsample_initial_channel, 7, 1, padding=3) + + self.ups = nn.ModuleList() + for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): + self.ups.append( + nn.ModuleList( + [ + ops.ConvTranspose1d( + upsample_initial_channel // (2 ** i), + upsample_initial_channel // (2 ** (i + 1)), + k, + u, + padding=(k - u) // 2, + ) + ] + ) + ) + + self.resblocks = nn.ModuleList() + for i in range(len(self.ups)): + ch = upsample_initial_channel // (2 ** (i + 1)) + for k, d in zip(resblock_kernel_sizes, resblock_dilation_sizes): + self.resblocks.append(AMPBlock1(ch, k, d)) + + self.activation_post = Activation1d(activation=SnakeBeta(ch)) + self.conv_post = ops.Conv1d(ch, 1, 7, 1, padding=3, bias=False) + + def forward(self, x): + x = self.conv_pre(x) + + for i in range(self.num_upsamples): + for i_up in range(len(self.ups[i])): + x = self.ups[i][i_up](x) + xs = None + for j in range(self.num_kernels): + if xs is None: + xs = self.resblocks[i * self.num_kernels + j](x) + else: + xs += self.resblocks[i * self.num_kernels + j](x) + x = xs.div_(self.num_kernels) + + x = self.activation_post(x) + return self.conv_post(x).clamp_(-1.0, 1.0) + + +# Top-level VAE + +class MiniMaxH3AudioVAE(nn.Module): + """MiniMax H3 stereo audio VAE at 32 kHz. + + Latents are [B, 32, 2, T]: 32 channels, 2 stereo channels, T frames at + 40 latent frames per second (800 audio samples per latent frame). The + stereo channels are processed independently by the mono encoder/decoder. + Latents are normalized with the stored per-channel latents_mean/std. + """ + + def __init__( + self, + encoder_dim=64, + encoder_rates=(2, 4, 4, 5, 5), + latent_dim=2048, + decoder_dim=1024, + vae_latent_channels=32, + ): + super().__init__() + self.sample_rate = 32000 + + self.hop_length = 1 + for r in encoder_rates: + self.hop_length *= r + self.samples_per_latent = self.hop_length # 800 + self.latents_per_second = self.sample_rate // self.hop_length # 40 + self.output_sample_rate = self.sample_rate # read by LTXVAudioVAEDecode + + self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) + + self.pre_block = AttnProjection(latent_dim, vae_latent_channels, num_heads=8) + + self.mean_proj = ops.Conv1d(vae_latent_channels, vae_latent_channels, 1) + # logs_proj exists in the checkpoint but is unused at inference + # (encode returns the posterior mean, no sampling). + self.logs_proj = ops.Conv1d(vae_latent_channels, vae_latent_channels, 1) + + self.dec_in_proj = ops.Conv1d(vae_latent_channels, latent_dim, 1) + self.decoder = BigVGAN(num_mels=latent_dim, upsample_initial_channel=decoder_dim) + + self.register_buffer("latents_mean", torch.empty(vae_latent_channels)) + self.register_buffer("latents_std", torch.empty(vae_latent_channels)) + + def decode(self, z): + """Decode normalized latents [B, 32, 2, T] to stereo waveforms [B, 2, L] at 32 kHz.""" + b, c, s, t = z.shape + z = z.permute(0, 2, 1, 3).reshape(b * s, c, t) + mean = self.latents_mean.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + std = self.latents_std.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + z = z * std + mean + x = self.dec_in_proj(z) + x = self.decoder(x) # [b * s, 1, L], already clamped to [-1, 1] + return x.reshape(b, s, -1) + + def encode(self, waveform): + """Encode stereo waveforms [B, 2, L] at 32 kHz (in [-1, 1]) to normalized latents [B, 32, 2, T]. + + L is right-padded with zeros to a multiple of 800 samples; the returned + posterior mean is used directly (no sampling). + """ + b, s, length = waveform.shape + right_pad = math.ceil(length / self.hop_length) * self.hop_length - length + waveform = F.pad(waveform, (0, right_pad)) + x = waveform.reshape(b * s, 1, -1) + x = self.encoder(x) # [b * s, latent_dim, T] + x = self.pre_block(x.transpose(1, 2)).transpose(1, 2) # [b * s, 32, T] + z = self.mean_proj(x) + mean = self.latents_mean.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + std = self.latents_std.view(1, -1, 1).to(device=z.device, dtype=z.dtype) + z = (z - mean) / std + return z.reshape(b, s, z.shape[1], z.shape[2]).permute(0, 2, 1, 3) diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py new file mode 100644 index 000000000..494350d40 --- /dev/null +++ b/comfy/ldm/minimax/model.py @@ -0,0 +1,646 @@ +"""MiniMax H3 audio-video DiT. + +Single-stream packed-token transformer denoising video (24ch, patch 1x2x2) and +stereo audio (32ch, 40 Hz) latents jointly, conditioned on Qwen3-VL layer-50 hidden states. +The packed sequence is: +[text | cond rows | audio | video] for t2va/fl2va +[text | reference blocks | audio | video] for ref2va + +Timestep domain: the model receives the *video* sigma from the sampler and +derives per-token timesteps t = 1 - sigma internally; the audio stream runs on +its own shifted schedule (sigma_shift video 12.0 / audio 3.0), mapped from the +video sigma in closed form. The audio velocity is returned scaled by the +schedule map's derivative d(sigma_a)/d(sigma_v). +""" + +import math + +import torch +import torch.nn as nn + +import comfy.ldm.common_dit +import comfy.model_management +import comfy.model_prefetch +import comfy.ops +import comfy.patcher_extension +import comfy.quant_ops +from comfy.ldm.modules.attention import optimized_attention + +FRAME_PER_TOKEN = (1, 4, 4, 4, 4) +FRAME_RESCALE = 5.0 / 3.0 +VISUAL_COND_TIMESTEP = 0.999 +AUDIO_COND_TIMESTEP = 1.0 + + +def time_shift_sigma(sigma, from_shift, to_shift): + # invert sigma = s*b/(1+(s-1)*b) to the base grid, re-apply the other shift + base = sigma / (from_shift + sigma * (1.0 - from_shift)) + return to_shift * base / (1.0 + (to_shift - 1.0) * base) + + +def time_shift_slope(sigma, from_shift, to_shift): + """d(sigma_to)/d(sigma_from) at the same base-grid point. + + Scaling a stream's returned velocity by this slope makes the flat ODE that + any sampler integrates on the from-schedule equal to that stream's true ODE + on its own schedule. + """ + base = sigma / (from_shift + sigma * (1.0 - from_shift)) + return (to_shift * (1.0 + (from_shift - 1.0) * base) ** 2) / (from_shift * (1.0 + (to_shift - 1.0) * base) ** 2) + + +def patchify_video(latent, patch_size=(1, 2, 2)): + # [B, C, T, H, W] -> [B*t*h*w, C*pt*ph*pw] + b, c, t_full, h_full, w_full = latent.shape + pt, ph, pw = patch_size + t, h, w = t_full // pt, h_full // ph, w_full // pw + x = latent.reshape(b, c, t, pt, h, ph, w, pw) + x = torch.einsum("nctrhpwq->nthwcrpq", x) + return x.reshape(b * t * h * w, c * pt * ph * pw) + + +def unpatchify_video(rows, t, h, w, c=24, patch_size=(1, 2, 2)): + pt, ph, pw = patch_size + x = rows.reshape(-1, t, h, w, c, pt, ph, pw) + x = torch.einsum("nthwcrpq->nctrhpwq", x) + return x.reshape(-1, c, t * pt, h * ph, w * pw) + + +def pack_audio(latent): + # [B, C=32, ch=2, T] -> [ch*T, 32] channel-major (ch0 t0..T-1, ch1 t0..T-1) + b, c, ch, t = latent.shape + return latent[0].permute(1, 2, 0).reshape(ch * t, c) + + +def unpack_audio(rows, ch=2): + t = rows.shape[0] // ch + return rows.reshape(ch, t, rows.shape[-1]).permute(2, 0, 1).unsqueeze(0) + + +def _axis_from_sqrt_area(dim, patch, sqrt_area): + # linspace((1 - ratio) / 2, (1 + ratio) / 2, dim // patch, endpoint=False) * 32 + ratio = dim / sqrt_area + n = dim // patch + return (torch.arange(n, dtype=torch.float64) * (ratio / n) + (1.0 - ratio) / 2.0) * 32.0 + + +def _frame_grid(h, w): + # area-normalized (h, w) coordinates of one latent frame's 2x2-patch rows + area = math.sqrt(h * w) + hh, ww = torch.meshgrid(_axis_from_sqrt_area(h, 2, area), _axis_from_sqrt_area(w, 2, area), indexing="ij") + return torch.stack([hh.reshape(-1), ww.reshape(-1)], dim=-1), _axis_from_sqrt_area(w, 2, area) + + +def _video_t_spans(n): + return [FRAME_RESCALE * FRAME_PER_TOKEN[k % 5] for k in range(n)] + + +def _video_t_grid(n, origin): + # origin + exclusive cumsum + spans = torch.tensor(_video_t_spans(n), dtype=torch.float64) + return float(origin) + torch.cat([torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(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) + g[:, 0] = (cursor + torch.arange(t, dtype=torch.float64)).repeat(2) + g[:t, 2] = w_low + g[t:, 2] = w_high + return g + + +def _video_grid(vt, frame, cursor): + g = torch.empty(vt, frame.shape[0], 3, dtype=torch.float64) + g[:, :, 0] = _video_t_grid(vt, cursor)[:, None] + g[:, :, 1:] = frame[None] + return g.reshape(-1, 3) + + +class TimeEmbedder(nn.Module): + def __init__(self, freq_dim, hidden, out, dtype=None, device=None, operations=None): + super().__init__() + self.freq_dim = freq_dim + self.proj_in = operations.Linear(freq_dim, hidden, bias=True, dtype=dtype, device=device) + self.proj_out = operations.Linear(hidden, out, bias=True, dtype=dtype, device=device) + + def forward(self, t): + # t: [M] in [0, 1]; fp32 throughout, cos before sin + half = self.freq_dim // 2 + freqs = torch.exp(-math.log(10000.0) * torch.arange(half, dtype=torch.float32, device=t.device) / half) + args = t.to(torch.float32)[:, None] * freqs[None] + emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + return self.proj_out(nn.functional.silu(self.proj_in(emb))) + + +def rope_rotation_table(angles, dtype): + """[S, rot_dim] pair angles -> [1, S, 1, rot_dim/2, 2, 2] rotation matrices.""" + half = angles.shape[-1] // 2 + ang = angles[:, :half] # duplicated halves: [:, :half] == [:, half:] + c, s = torch.cos(ang), torch.sin(ang) + table = torch.stack([c, -s, s, c], dim=-1).reshape(1, angles.shape[0], 1, half, 2, 2) + return table.to(dtype) + + +class Attention(nn.Module): + def __init__(self, hidden, heads, head_dim, eps, dtype=None, device=None, operations=None): + super().__init__() + self.heads = heads + self.head_dim = head_dim + inner = heads * head_dim + self.qkv_proj = operations.Linear(hidden, inner * 3, bias=False, dtype=dtype, device=device) + self.q_norm = operations.RMSNorm(head_dim, eps=eps, dtype=dtype, device=device) + self.k_norm = operations.RMSNorm(head_dim, eps=eps, dtype=dtype, device=device) + self.out_proj = operations.Linear(inner, hidden, bias=False, dtype=dtype, device=device) + + def forward(self, x, rope_freqs=None, transformer_options={}): + s = x.shape[0] + q, k, v = self.qkv_proj(x).split(self.heads * self.head_dim, dim=-1) + v = v.view(s, self.heads, self.head_dim) + if rope_freqs is not None: + # fused per-head RMSNorm + partial split-half rope, in place on the qkv buffer + q = q.view(1, s, self.heads, self.head_dim) + k = k.view(1, s, self.heads, self.head_dim) + qw = comfy.model_management.cast_to(self.q_norm.weight, device=x.device) + kw = comfy.model_management.cast_to(self.k_norm.weight, device=x.device) + rot = rope_freqs.shape[-3] * 2 + if comfy.model_management.in_training: + q, k = comfy.quant_ops.ck.rms_rope_split_half( + q, k, rope_freqs, qw, kw, epsilon=self.q_norm.eps, rot_dim=rot) + else: + comfy.quant_ops.ck.rms_rope_split_half_( + q, k, rope_freqs, qw, kw, epsilon=self.q_norm.eps, rot_dim=rot) + q = q[0] + k = k[0] + 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) + out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) + return self.out_proj(out.squeeze(0)) + + +class MLP(nn.Module): + def __init__(self, hidden, ffn, dtype=None, device=None, operations=None): + super().__init__() + self.fc1 = operations.Linear(hidden, ffn * 2, bias=False, dtype=dtype, device=device) + self.fc2 = operations.Linear(ffn, hidden, bias=False, dtype=dtype, device=device) + + def forward(self, x): + return comfy.ops.linear_input_act(self.fc2, self.fc1(x), "swiglu") + + +class AdalnProj(nn.Module): + def __init__(self, t_dim, hidden, expand, modalities, apply_silu=True, + dtype=None, device=None, operations=None): + super().__init__() + self.expand = expand + self.modalities = modalities + self.hidden = hidden + self.apply_silu = apply_silu + self.linear = operations.Linear(t_dim, expand * hidden * modalities, bias=True, dtype=dtype, device=device) + + def forward(self, t_emb): + # [M, t_dim] -> expand tensors of [M*modalities, hidden] + x = self.linear(nn.functional.silu(t_emb) if self.apply_silu else t_emb) + x = x.view(x.shape[0] * self.modalities, self.expand * self.hidden) + return x.chunk(self.expand, dim=-1) + + +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)) + 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)) + return x + + +class RefinerBlock(nn.Module): + def __init__(self, hidden, heads, head_dim, ffn, eps, qk_eps, dtype=None, device=None, operations=None): + super().__init__() + self.norm1 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.norm2 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.attn = Attention(hidden, heads, head_dim, qk_eps, dtype=dtype, device=device, operations=operations) + self.mlp = MLP(hidden, ffn, dtype=dtype, device=device, operations=operations) + + def forward(self, x, transformer_options={}): + # attn/mlp outputs are fresh: accumulate residuals in place + x = self.attn(self.norm1(x), transformer_options=transformer_options).add_(x) + return self.mlp(self.norm2(x)).add_(x) + + +class TokenRefiner(nn.Module): + def __init__(self, num_layers, hidden, heads, head_dim, ffn, eps, qk_eps, final_eps, + dtype=None, device=None, operations=None): + super().__init__() + self.blocks = nn.ModuleList([ + RefinerBlock(hidden, heads, head_dim, ffn, eps, qk_eps, dtype=dtype, device=device, operations=operations) + for _ in range(num_layers)]) + self.final_norm = operations.RMSNorm(hidden, eps=final_eps, dtype=dtype, device=device) + + def forward(self, x, transformer_options={}): + for block in self.blocks: + x = block(x, transformer_options=transformer_options) + return self.final_norm(x) + + +class DiTBlock(nn.Module): + def __init__(self, hidden, heads, head_dim, ffn, t_dim, eps, qk_eps, + apply_silu=True, adaln_dtype=None, dtype=None, device=None, operations=None): + super().__init__() + self.norm1 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.norm2 = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.attn = Attention(hidden, heads, head_dim, qk_eps, dtype=dtype, device=device, operations=operations) + self.mlp = MLP(hidden, ffn, dtype=dtype, device=device, operations=operations) + self.adaln_proj = AdalnProj(t_dim, hidden, 6, 3, apply_silu=apply_silu, + dtype=adaln_dtype if adaln_dtype is not None else dtype, + device=device, operations=operations) + + def forward(self, x, t_emb, mod_segments, rope_freqs, transformer_options={}): + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaln_proj(t_emb) + h = _mod_scale_shift(self.norm1(x), shift_msa, scale_msa, mod_segments) + x = _mod_gate(x, gate_msa, self.attn(h, rope_freqs=rope_freqs, transformer_options=transformer_options), mod_segments) + h = _mod_scale_shift(self.norm2(x), shift_mlp, scale_mlp, mod_segments) + return _mod_gate(x, gate_mlp, self.mlp(h), mod_segments) + + +class FinalLayer(nn.Module): + def __init__(self, hidden, t_dim, video_dim, audio_dim, eps, apply_silu=True, adaln_dtype=None, + dtype=None, device=None, operations=None): + super().__init__() + self.norm = operations.RMSNorm(hidden, eps=eps, dtype=dtype, device=device) + self.adaln_proj = AdalnProj(t_dim, hidden, 2, 1, apply_silu=apply_silu, + dtype=adaln_dtype if adaln_dtype is not None else dtype, + device=device, operations=operations) + # output heads are the checkpoint's fp32 island; norm/adaln are stored at model dtype + self.video_out = operations.Linear(hidden, video_dim, bias=True, dtype=torch.float32, device=device) + 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 + 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) + + +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): + frame, w_grid = _frame_grid(latent_h, latent_w) + frame_rows = frame.shape[0] + + segments = [("text", text_len)] # (kind, n_rows) + g = torch.zeros(text_len, 3, dtype=torch.float64) + g[:, 0] = torch.arange(text_len, dtype=torch.float64) + pos = [g] # per segment: [n, 3] float64 (t, h, w) + + 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])) + if refs: + cursor = float(text_len) + for blk in refs: + kind = blk["kind"] + if kind == "image": + r_frame, _ = _frame_grid(blk["latent_h"], blk["latent_w"]) + n = r_frame.shape[0] + g = torch.empty(n, 3, dtype=torch.float64) + g[:, 0] = cursor + g[:, 1:] = r_frame + segments.append(("ref_img", n)) + pos.append(g) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + cursor += 1.0 + elif kind == "audio": + rt = blk["ref_audio_t"] + if rt > 0: + segments.append(("ref_audio", rt * 2)) + pos.append(_audio_grid(cursor, 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 + cursor += float(rt) + elif kind in ("video", "video_audio"): + # the block's audio rows pack immediately before its video + # rows, both sharing the cursor origin + rt = blk["ref_audio_t"] + vt = blk["latent_t"] + r_frame, r_w_grid = _frame_grid(blk["latent_h"], blk["latent_w"]) + if rt > 0: + segments.append(("ref_audio", rt * 2)) + pos.append(_audio_grid(cursor, rt, float(r_w_grid[0]), float(r_w_grid[-1]))) + audio_pos.append(torch.arange(row, row + rt * 2)) + audio_update.append(torch.zeros(rt * 2, dtype=torch.bool)) + row += rt * 2 + n = vt * r_frame.shape[0] + segments.append(("ref_img", n)) + pos.append(_video_grid(vt, r_frame, cursor)) + img_pos.append(torch.arange(row, row + n)) + img_update.append(torch.zeros(n, dtype=torch.bool)) + row += n + cursor += max(float(rt), sum(_video_t_spans(vt))) + + # target audio then target video, always the last two segments + segments.append(("audio", audio_t * 2)) + pos.append(_audio_grid(cursor, audio_t, *target_audio_w)) + audio_pos.append(torch.arange(row, row + audio_t * 2)) + audio_update.append(torch.ones(audio_t * 2, dtype=torch.bool)) + row += audio_t * 2 + + n_video = latent_t * frame_rows + segments.append(("video", n_video)) + pos.append(_video_grid(latent_t, frame, cursor)) + img_pos.append(torch.arange(row, row + n_video)) + img_update.append(torch.ones(n_video, dtype=torch.bool)) + row += n_video + + self.seq_len = row + self.position_ids = torch.cat(pos) # [S, 3] float64 + self.img_pos = torch.cat(img_pos) + self.img_update = torch.cat(img_update) + self.audio_pos = torch.cat(audio_pos) + 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 + # 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 = [] + off = 0 + for kind, n in segments: + seg_abs.append((off, off + n, kind)) + off += n + self.segments = seg_abs + + +class MiniMaxH3Model(nn.Module): + def __init__(self, hidden_size=5376, num_layers=50, token_refiner_num_layers=2, + num_attention_heads=56, attention_head_dim=128, ffn_hidden_size=14336, + latents_dim=24, audio_latents_dim=32, patch_size=(1, 2, 2), text_dim=5120, + timestep_input_dim=256, time_embed_hidden_size=5376, time_embed_dim=2688, + rope_inv_freq_len=16, norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5, + sigma_shift_video=12.0, sigma_shift_audio=3.0, + adaln_curve_grid=None, + image_model=None, dtype=None, device=None, operations=None, **kwargs): + super().__init__() + self.dtype = dtype + self.hidden_size = hidden_size + self.patch_size = tuple(patch_size) + self.latents_dim = latents_dim + self.audio_latents_dim = audio_latents_dim + self.sigma_shift_video = sigma_shift_video + self.sigma_shift_audio = sigma_shift_audio + self.use_adaln_curves = adaln_curve_grid is not None + # curve-form checkpoints replace the time embedder and full-width adaln weights with a small shared basis of the time-embedding curve + curve = {"apply_silu": not self.use_adaln_curves, + "adaln_dtype": torch.float32 if self.use_adaln_curves else dtype} + video_patch_dim = latents_dim * self.patch_size[0] * self.patch_size[1] * self.patch_size[2] + + self.video_patch_proj = operations.Linear(video_patch_dim, hidden_size, bias=True, dtype=torch.float32, device=device) + self.audio_patch_proj = operations.Linear(audio_latents_dim, hidden_size, bias=True, dtype=torch.float32, device=device) + self.condition_proj = operations.Linear(text_dim, hidden_size, bias=True, dtype=dtype, device=device) + if self.use_adaln_curves: + self.register_buffer("adaln_t_table", torch.empty(adaln_curve_grid, time_embed_dim, dtype=torch.float32)) + else: + self.time_embedder = TimeEmbedder(timestep_input_dim, time_embed_hidden_size, time_embed_dim, + dtype=torch.float32, device=device, operations=operations) + self.rope = nn.Module() + self.rope.register_buffer("inv_freq", torch.empty(rope_inv_freq_len, dtype=torch.float32)) + self.token_refiner = TokenRefiner(token_refiner_num_layers, hidden_size, num_attention_heads, + attention_head_dim, ffn_hidden_size, norm_eps, qk_norm_eps, + final_norm_eps, dtype=dtype, device=device, operations=operations) + self.blocks = nn.ModuleList([ + DiTBlock(hidden_size, num_attention_heads, attention_head_dim, ffn_hidden_size, + time_embed_dim, norm_eps, qk_norm_eps, **curve, dtype=dtype, device=device, operations=operations) + for _ in range(num_layers)]) + self.final_layer = FinalLayer(hidden_size, time_embed_dim, video_patch_dim, audio_latents_dim, + final_norm_eps, **curve, dtype=dtype, device=device, operations=operations) + + def preprocess_text_embeds(self, text_states): + """[B, L, text_dim] Qwen states -> [B, L, hidden] refined text embeds.""" + if text_states.shape[-1] == self.hidden_size: + return text_states + return self.token_refiner(self.condition_proj(text_states[0])).unsqueeze(0) + + def rope_freqs(self, position_ids, device): + # [S, 3] float64 -> [S, 96] fp32 + pos = position_ids.to(torch.float32).to(device) + inv = comfy.model_management.cast_to(self.rope.inv_freq, device=device) + per_axis = pos.unsqueeze(-1) * inv.view(1, 1, -1) # [S, 3, 16] + t_f, h_f, w_f = per_axis.unbind(dim=1) + half = torch.cat((t_f, h_f, w_f), dim=-1) # [S, 48] + return torch.cat((half, half), dim=-1) # [S, 96] + + def _cond_video_rows(self, payload, device): + """Concatenated visual condition rows (normalized latents -> patchified), with condition noise augmentation.""" + rows = [] + aug = payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP) + seed = int(payload.get("seed", 0)) + # every condition intentionally restarts the same RNG stream + for z in payload.get("cond_video_latents", []): + r = patchify_video(z.to(torch.float32), self.patch_size) + if aug < 1.0: + gen = torch.Generator("cpu").manual_seed(seed) + noise = torch.randn(r.shape, generator=gen, dtype=torch.float32) + r = aug * r + (1.0 - aug) * noise.to(r.device) + rows.append(r.to(device)) + return torch.cat(rows, dim=0) if rows else None + + def _cond_audio_rows(self, payload, device): + rows = [] + aug = payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP) + seed = int(payload.get("seed", 0)) + 1 + for z in payload.get("cond_audio_latents", []): + r = pack_audio(z.to(torch.float32)) + if aug < 1.0: + gen = torch.Generator("cpu").manual_seed(seed) + noise = torch.randn(r.shape, generator=gen, dtype=torch.float32) + r = aug * r + (1.0 - aug) * noise.to(r.device) + 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): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload, **kwargs) + + def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=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) + if video_x.shape[0] != 1: + raise ValueError("MiniMax H3 supports batch size 1") + payload = minimax_payload or {} + device = video_x.device + dtype = context.dtype # compute dtype + + latent_t, lat_h, lat_w = video_x.shape[2], video_x.shape[3], video_x.shape[4] + audio_t = audio_x.shape[-1] + text_len = context.shape[1] + # extra_conds prebuilds the layout once per sampling run + layout = payload.get("layout") + 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")) + + # model_base passes model_sampling.timestep(sigma) = sigma * 1000 + shift_v = float(transformer_options.get("minimax_h3_sigma_shift_video", self.sigma_shift_video)) + shift_a = float(transformer_options.get("minimax_h3_sigma_shift_audio", self.sigma_shift_audio)) + sigma_v = (timestep.flatten()[0] / 1000.0).float().clamp(min=1e-6) + t_v = float(1.0 - sigma_v) + t_a = float(1.0 - time_shift_sigma(sigma_v, shift_v, shift_a)) + + # 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())) + 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} + + text_tags = payload.get("text_token_tags") + mod_segments = [] + for a, b, kind in layout.segments: + row_base = t_row[seg_t[kind]] * 3 + if kind == "text" and text_tags is not None: + # the presentation text span mixes tags (vision pads carry the video modality) split into tag runs + tags = text_tags.view(-1).tolist() + run_start = 0 + for i in range(1, b - a + 1): + 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 + else: + mod_segments.append((a, b, row_base + seg_tag[kind])) + + # embed + img_update = layout.img_update.to(device) + audio_update = layout.audio_update.to(device) + video_rows = patchify_video(video_x.to(torch.float32), self.patch_size) + audio_rows = pack_audio(audio_x.to(torch.float32)) + cond_video_rows = self._cond_video_rows(payload, device) + cond_audio_rows = self._cond_audio_rows(payload, device) + + all_video_rows = video_rows + if cond_video_rows is not None: + all_video_rows = torch.empty(img_update.shape[0], video_rows.shape[1], dtype=torch.float32, device=device) + all_video_rows[~img_update] = cond_video_rows + all_video_rows[img_update] = video_rows + all_audio_rows = audio_rows + if cond_audio_rows is not None: + all_audio_rows = torch.empty(audio_update.shape[0], audio_rows.shape[1], dtype=torch.float32, device=device) + all_audio_rows[~audio_update] = cond_audio_rows + all_audio_rows[audio_update] = audio_rows + + video_embed = self.video_patch_proj(all_video_rows).to(dtype) + audio_embed = self.audio_patch_proj(all_audio_rows).to(dtype) + text_states = context[0] + if text_states.shape[-1] != self.hidden_size: + text_states = self.token_refiner(self.condition_proj(text_states), + transformer_options=transformer_options) + + # segments are contiguous: assemble by slices, embed rows follow segment order + h = torch.empty(layout.seq_len, self.hidden_size, dtype=dtype, device=device) + voff = aoff = 0 + for a, b, kind in layout.segments: + n = b - a + if kind == "text": + h[a:b] = text_states + elif kind in ("cond", "ref_img", "video"): + h[a:b] = video_embed[voff:voff + n] + voff += n + else: # ref_audio / audio + h[a:b] = audio_embed[aoff:aoff + n] + aoff += n + + t_vals = torch.tensor(unique_t, dtype=torch.float32, device=device) + if self.use_adaln_curves: + # adaln projections consume interpolated coordinates of the time-embedding curve + table = comfy.model_management.cast_to(self.adaln_t_table, device=device) + pos = t_vals.clamp(0.0, 1.0) * (table.shape[0] - 1) # t in [0,1] -> fractional grid index, out-of-range t clamps to the curve ends + i0 = pos.floor().long().clamp(max=table.shape[0] - 2) # lower grid row, max-clamp keeps t=1.0 on the last interval instead of reading past the table + t_emb = torch.lerp(table[i0], table[i0 + 1], (pos - i0).unsqueeze(1)) # blend the two rows by the fractional part + else: + t_emb = self.time_embedder(t_vals).to(dtype) + + # rotation table computed once per forward, consumed by the kitchen split-half rope + rope_freqs = rope_rotation_table(self.rope_freqs(layout.position_ids, device), dtype) + + # blocks + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + prefetch_queue = comfy.model_prefetch.make_prefetch_queue(list(self.blocks), device, transformer_options) + for i, block in enumerate(self.blocks): + comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, block) + if ("double_block", i) in blocks_replace: + def block_wrap(args): + return {"img": block(args["img"], args["t_emb"], args["mod_segments"], args["rope_freqs"], + transformer_options=args["transformer_options"])} + h = blocks_replace[("double_block", i)]( + {"img": h, "t_emb": t_emb, "mod_segments": mod_segments, "rope_freqs": rope_freqs, + "transformer_options": transformer_options}, + {"original_block": block_wrap})["img"] + else: + h = block(h, t_emb, mod_segments, rope_freqs, transformer_options=transformer_options) + if prefetch_queue is not None: + 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") + 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) + video_out = video_out[:, :, :orig_t, :orig_h, :orig_w] + audio_out = unpack_audio(a) + + # The sampler integrates the flat ODE dX/dsigma_v = (X - denoised)/sigma_v. + # Scaling the audio velocity by d(sigma_a)/d(sigma_v) makes that ODE equal + # to the audio stream's true ODE on its own shifted schedule. + slope_a = time_shift_slope(sigma_v, shift_v, shift_a).to(audio_out.dtype) + return [-video_out.to(video_x.dtype), (-slope_a) * audio_out.to(audio_x.dtype)] diff --git a/comfy/ldm/minimax/vae.py b/comfy/ldm/minimax/vae.py new file mode 100644 index 000000000..aeb3421a2 --- /dev/null +++ b/comfy/ldm/minimax/vae.py @@ -0,0 +1,694 @@ +# MiniMax H3 video VAE: 3D causal CNN encoder + ViT3D decoder. + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import comfy.ops +import comfy.quant_ops +import comfy.rmsnorm +from comfy.ldm.modules.attention import optimized_attention + +ops = comfy.ops.disable_weight_init + +IMAGENET_MEAN = (0.485, 0.456, 0.406) +IMAGENET_STD = (0.229, 0.224, 0.225) + +LATENTS_MEAN = [ + 0.858090341091156, -0.9606591463088989, 1.0661640167236328, -0.5090325474739075, + -0.2727581858634949, -1.3675414323806763, -0.2553254961967468, -0.26907554268836975, + -0.5376840829849243, -0.0464097298681736, 0.6657370328903198, 0.19690127670764923, + -0.5460608005523682, -0.4035342037677765, -0.23683024942874908, 0.25928452610969543, + -0.30133944749832153, 0.211341992020607, -1.1206848621368408, 0.3581933379173279, + -0.04225143790245056, 0.2604829967021942, 0.22864092886447906, 0.7056031823158264, +] + +LATENTS_STD = [ + 1.2223774194717407, 1.2767263650894165, 1.68317747116088865, 1.7549455165863037, + 1.5636216402053833, 2.194143533706665, 0.96531379222869875, 1.05698859691619875, + 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, + 0.7996809482574463, 0.44988900423049925, 0.7197399735450745, 0.69362932443618775, + 2.961095094680786, 2.7694199085235595, 3.0496184825897215, 2.1088054180145265, + 3.276226282119751, 3.1627357006073, 2.28168129920959475, 2.6127843856811525, +] + + +# 3D causal CNN encoder + +class CausalConv3d(ops.Conv3d): + # Reflect spatial padding, causal (zeros, front-only) temporal padding. + def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0): + super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride) + self.causal_padding = (padding,) * 3 if isinstance(padding, int) else tuple(padding) + + def forward(self, x): + if sum(self.causal_padding) == 0: + return super().forward(x) + + x = F.pad(x, (self.causal_padding[2], self.causal_padding[2], self.causal_padding[1], self.causal_padding[1], 0, 0), mode="reflect") + if x.shape[2] == 1: + # single frame: the causal front padding is all zeros truncate the temporal taps instead of convolving zero frames + return super().forward(x, autopad="causal_zero") + x = F.pad(x, (0, 0, 0, 0, self.causal_padding[0] * 2, 0), mode="constant") + return super().forward(x) + + +class TemporalIsolatedGroupNorm(ops.GroupNorm): + # GroupNorm with statistics computed per frame (time merged into batch). + def forward(self, x): + if x.dim() == 5: + b, c, t, h, w = x.shape + x = x.permute(0, 2, 1, 3, 4).contiguous().view(b * t, c, 1, h, w) + x = super().forward(x) + return x.view(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + return super().forward(x) + + +def group_norm_3d(num_channels): + return TemporalIsolatedGroupNorm(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True) + + +class Downsample3D(nn.Module): + def __init__(self, in_channels, out_channels, time_stride=1, space_stride=2): + super().__init__() + self.space_stride = space_stride + self.conv = CausalConv3d( + in_channels, + out_channels, + kernel_size=3, + padding=(1, 0, 0), + stride=(time_stride, space_stride, space_stride), + ) + + def forward(self, x): + if self.space_stride == 2: + x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect") + return self.conv(x) + + +class ResnetBlock3D(nn.Module): + def __init__(self, in_channels, out_channels=None): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + + self.norm1 = group_norm_3d(in_channels) + self.norm2 = group_norm_3d(out_channels) + self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, padding=1) + self.conv2 = CausalConv3d(out_channels, out_channels, kernel_size=3, padding=1) + if in_channels != out_channels: + self.nin_shortcut = CausalConv3d(in_channels, out_channels, kernel_size=1) + + def forward(self, x): + h = self.conv1(F.silu(self.norm1(x), inplace=True)) + h = self.conv2(F.silu(self.norm2(h), inplace=True)) + if self.in_channels != self.out_channels: + x = self.nin_shortcut(x) + return h.add_(x) + + +class EncoderFCN3D(nn.Module): + def __init__(self, ch, ch_mult, space_down, time_down, num_res_blocks, in_channels, z_channels, double_z=True): + super().__init__() + self.num_levels = len(ch_mult) + if isinstance(num_res_blocks, int): + num_res_blocks = [num_res_blocks] * self.num_levels + self.num_res_blocks = num_res_blocks + + block_mid = [ch * ch_mult[i] for i in range(self.num_levels)] + block_in = [block_mid[0]] + block_mid[:-1] + block_out = block_mid + + self.conv_in = CausalConv3d(in_channels, block_in[0], kernel_size=3, padding=1) + + self.down = nn.ModuleList() + for i_level in range(self.num_levels): + down = nn.Module() + down.block = nn.ModuleList() + for i in range(self.num_res_blocks[i_level]): + down.block.append( + ResnetBlock3D( + in_channels=block_in[i_level] if i == 0 else block_mid[i_level], + out_channels=block_mid[i_level], + ) + ) + if space_down[i_level] * time_down[i_level] > 1: + down.downsample = Downsample3D( + block_mid[i_level], + block_out[i_level], + time_stride=time_down[i_level], + space_stride=space_down[i_level], + ) + self.down.append(down) + + self.norm_out = group_norm_3d(block_out[-1]) + self.conv_out = CausalConv3d( + block_out[-1], + 2 * z_channels if double_z else z_channels, + kernel_size=3, + padding=1, + ) + + def forward(self, x): + h = self.conv_in(x) + for i_level in range(self.num_levels): + for i_block in range(self.num_res_blocks[i_level]): + h = self.down[i_level].block[i_block](h) + if hasattr(self.down[i_level], "downsample"): + h = self.down[i_level].downsample(h) + h = F.silu(self.norm_out(h)) + return self.conv_out(h) + + +# ViT3D decoder + +def create_token_ids(patch_dims, device, dtype): + coords_list = [] + for dim_size in patch_dims: + coords = torch.arange(0.5, dim_size, dtype=dtype, device=device) + coords = coords / dim_size + coords = 2.0 * coords - 1.0 + coords_list.append(coords) + coords = torch.stack(torch.meshgrid(*coords_list, indexing="ij"), dim=-1) + return coords.flatten(0, len(patch_dims) - 1).unsqueeze(0) + + +class RotaryEmbeddingND(nn.Module): + def __init__(self, dim, rotary_base=100.0, n_dim=3): + super().__init__() + self.n_dim = n_dim + self.angle_scale = 2.0 * math.pi + inv_freq = 1 / rotary_base ** torch.arange(0, 1, 2 * n_dim / dim, dtype=torch.float32) + self.register_buffer("inv_freq", inv_freq, persistent=False) + + def forward(self, img_ids): + # [B, S, n_dim] -> [B, S, 1, pairs, 2, 2] rotation table for the kitchen split-half rope + angles = ( + self.angle_scale + * img_ids[:, :, :, None].float() + * self.inv_freq.to(img_ids.device)[None, None, None, :] + ) + angles = angles.flatten(2, 3) + c, s = torch.cos(angles), torch.sin(angles) + table = torch.stack([c, -s, s, c], dim=-1).reshape(*angles.shape[:2], 1, angles.shape[-1], 2, 2) + return table.to(img_ids.dtype) + + +class FeedForward(nn.Module): + # Gated SiLU FFN. + def __init__(self, dim, mult=4, bias=True): + super().__init__() + inner_dim = dim * mult + self.w1 = ops.Linear(dim, inner_dim * 2, bias=bias) + self.w2 = ops.Linear(inner_dim, dim, bias=bias) + + def forward(self, x): + gate, x = self.w1(x).chunk(2, dim=-1) + return self.w2(F.silu(gate).mul_(x)) + + +class Attention(nn.Module): + def __init__(self, heads, dim_head, bias=True, eps=1e-5): + super().__init__() + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + self.norm_q = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) + self.norm_k = ops.RMSNorm(dim_head, eps=eps, elementwise_affine=False) + self.to_qkv = ops.Linear(inner_dim, inner_dim * 3, bias=bias) + self.to_out = ops.Linear(inner_dim, inner_dim, bias=bias) + + def forward(self, x, rotary_pos_emb=None): + batch_size, seq_len, _ = x.shape + + qkv = self.to_qkv(x) + qkv = qkv.view(batch_size, seq_len, -1, 3 * self.dim_head) + query, key, value = torch.chunk(qkv, 3, dim=-1) + + query = comfy.rmsnorm.rms_norm(query, self.norm_q.weight, self.norm_q.eps) + key = comfy.rmsnorm.rms_norm(key, self.norm_k.weight, self.norm_k.eps) + + if rotary_pos_emb is not None: + rot = rotary_pos_emb.shape[-3] * 2 + query[..., :rot], key[..., :rot] = comfy.quant_ops.ck.apply_rope_split_half( + query[..., :rot], key[..., :rot], rotary_pos_emb) + + out = optimized_attention(query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), + self.heads, skip_reshape=True).nan_to_num_(0.0) + return self.to_out(out) + + +class TransformerBlock(nn.Module): + def __init__(self, heads, dim_head, bias=True, eps=1e-5): + super().__init__() + dim = heads * dim_head + self.norm1 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) + self.attn = Attention(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + self.scale1 = nn.Parameter(torch.empty(dim)) + self.norm2 = ops.RMSNorm(dim, elementwise_affine=True, eps=eps) + self.ff = FeedForward(dim=dim, bias=bias) + self.scale2 = nn.Parameter(torch.empty(dim)) + + def forward(self, x, rotary_pos_emb=None): + x = x.addcmul_(self.attn(comfy.rmsnorm.rms_norm(x, self.norm1.weight, self.norm1.eps), rotary_pos_emb), self.scale1) + return x.addcmul_(self.ff(comfy.rmsnorm.rms_norm(x, self.norm2.weight, self.norm2.eps)), self.scale2) + + +class ViT3DDecoder(nn.Module): + def __init__(self, patch_size=16, patch_size_t=4, in_channels=24, out_channels=3, num_layers=36, heads=32, dim_head=64, rope_theta=100.0, + rope_dim_ratio=0.75, bias=True, eps=1e-5, num_register_tokens=4): + super().__init__() + dim = heads * dim_head + self.patch_size = patch_size + self.patch_size_t = patch_size_t + self.out_channels = out_channels + self.num_register_tokens = num_register_tokens + + self.pos_embed = RotaryEmbeddingND(int(dim_head * rope_dim_ratio), rope_theta, n_dim=3) + self.x_embedder = ops.Linear(in_channels, dim) + self.register_tokens = nn.Parameter(torch.empty(1, num_register_tokens, dim)) + # unused at inference; kept so the checkpoint loads without leftover keys + self.register_buffer("mask_token", torch.empty(1, 1, dim)) + + self.transformer_blocks = nn.ModuleList( + [TransformerBlock(heads=heads, dim_head=dim_head, bias=bias, eps=eps) + for _ in range(num_layers)] + ) + + self.norm_out = ops.LayerNorm(dim, elementwise_affine=True, eps=eps) + self.proj_out = ops.Linear(dim, out_channels * patch_size_t * patch_size * patch_size) + + def forward(self, x): + B, C, latent_T, latent_H, latent_W = x.shape + + h = self.x_embedder(x.flatten(2).transpose(1, 2)) # [B, T*H*W, C] + + num_patches = h.shape[1] + num_suffix = 1 + self.num_register_tokens + + h = torch.cat([h, self.register_tokens.expand(B, -1, -1), torch.zeros_like(h[:, 0:1, :])], dim=1) + + img_ids = create_token_ids((latent_T, latent_H, latent_W), x.device, x.dtype).expand(B, -1, -1) + suffix_ids = torch.zeros((B, num_suffix, 3), device=x.device, dtype=img_ids.dtype) + img_ids = torch.cat([img_ids, suffix_ids], dim=1) + + rotary_pos_emb = self.pos_embed(img_ids) + + for block in self.transformer_blocks: + h = block(h, rotary_pos_emb) + + output = self.proj_out(self.norm_out(h)) + + output = output[:, :num_patches, :] + + output = output.view( + B, latent_T, latent_H, latent_W, + self.out_channels, self.patch_size_t, self.patch_size, self.patch_size, + ) + output = output.permute(0, 4, 1, 5, 2, 6, 3, 7).contiguous() + output = output.reshape( + B, self.out_channels, + latent_T * self.patch_size_t, + latent_H * self.patch_size, + latent_W * self.patch_size, + ) + return output + + +# Full VAE + +class MiniMaxH3VideoVAE(nn.Module): + def __init__( + self, + in_channels=3, + out_ch=3, + ch=128, + embed_dim=24, + z_channels=24, + ch_mult=(1, 2, 2, 4, 4, 8), + num_res_blocks=2, + space_down=(2, 2, 2, 2, 1, 1), + time_down=(1, 2, 2, 1, 1, 1), + clip_length=17, + token_drop=3, + tile_size=256, + tile_overlap_min=64, + tiling=True, + ): + super().__init__() + self.vae_ratio = int(math.prod(space_down)) + self.vae_ratio_t = int(math.prod(time_down)) + + # temporal chunking parameters + self.clip_length = clip_length + self.token_drop = token_drop + self.frame_pre_padding = (-clip_length) % self.vae_ratio_t + self.tokens_chunk_size = math.ceil(clip_length / self.vae_ratio_t) + self.token_overlap = (-token_drop) % self.tokens_chunk_size + self.frame_overlap = max(self.token_overlap * self.vae_ratio_t - self.frame_pre_padding, 0) + + # spatial tiling parameters + self.tiling = tiling + self.tile_size = tile_size + self.tile_overlap_min = tile_overlap_min + + self.encoder = EncoderFCN3D( + ch=ch, + ch_mult=list(ch_mult), + space_down=list(space_down), + time_down=list(time_down), + num_res_blocks=num_res_blocks, + in_channels=in_channels, + z_channels=z_channels, + double_z=True, + ) + self.quant_conv = ops.Conv3d(z_channels * 2, 2 * embed_dim, 1) + self.post_quant_conv = ops.Conv3d(embed_dim, z_channels, 1) + self.decoder = ViT3DDecoder( + patch_size=self.vae_ratio, + patch_size_t=self.vae_ratio_t, + in_channels=z_channels, + out_channels=out_ch, + ) + + self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN)) + self.register_buffer("latents_std", torch.tensor(LATENTS_STD)) + self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1, 1), persistent=False) + self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False) + + # single-shot forward + + def _encode_moments(self, x): + return self.quant_conv(self.encoder(x)) + + def _decode_pixels(self, z): + return self.decoder(self.post_quant_conv(z)) + + def _adaptive_encode(self, x): + if self.tiling: + return self.tiled_encode(x) + return self._encode_moments(x) + + def _adaptive_decode(self, z): + if self.tiling: + return self.tiled_decode(z) + return self._decode_pixels(z) + + # spatial tiling + + def split_tiles(self, input_len): + tile_size = self.tile_size + if tile_size >= input_len: + return [0], [input_len], [] + + N = math.ceil(input_len / tile_size) + while True: + overlaps = [self.tile_overlap_min] * (N - 1) + remaining = tile_size * N - sum(overlaps) - input_len + if remaining < 0: + N += 1 + else: + break + + remaining_units = remaining // self.vae_ratio + for i in range(remaining_units): + overlaps[i % (N - 1)] += self.vae_ratio + + tile_start_idx = [0] + for i in range(N - 1): + tile_start_idx.append(tile_start_idx[-1] + tile_size - overlaps[i]) + + return tile_start_idx, [tile_size] * N, overlaps + + def blend(self, a, b, blend_extent, dim): + blend_extent = min(a.shape[dim], b.shape[dim], blend_extent) + + positions = torch.arange(blend_extent, device=b.device, dtype=b.dtype) + weight_a = 1 - positions / blend_extent + weight_b = positions / blend_extent + + shape = [1] * a.ndim + shape[dim] = blend_extent + weight_a = weight_a.view(shape) + weight_b = weight_b.view(shape) + + slice_a = [slice(None)] * a.ndim + slice_a[dim] = slice(-blend_extent, None) + slice_b = [slice(None)] * b.ndim + slice_b[dim] = slice(0, blend_extent) + + blended = a[tuple(slice_a)] * weight_a + b[tuple(slice_b)] * weight_b + + if blend_extent < b.shape[dim]: + slice_b_rest = [slice(None)] * b.ndim + slice_b_rest[dim] = slice(blend_extent, None) + return torch.cat([blended, b[tuple(slice_b_rest)]], dim=dim) + return blended + + def tiled_encode(self, x): + height, width = x.shape[-2], x.shape[-1] + y_idx, y_len, y_overlap = self.split_tiles(height) + x_idx, x_len, x_overlap = self.split_tiles(width) + + rows = [] + for i_pos, i_len in zip(y_idx, y_len): + row = [] + for j_pos, j_len in zip(x_idx, x_len): + tile = x[..., i_pos:i_pos + i_len, j_pos:j_pos + j_len] + row.append(self._encode_moments(tile)) + rows.append(row) + + latent_y_overlap = [o // self.vae_ratio for o in y_overlap] + latent_x_overlap = [o // self.vae_ratio for o in x_overlap] + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2) + if j > 0: + tile = self.blend(row[j - 1], tile, latent_x_overlap[j - 1], dim=-1) + if i < len(rows) - 1: + tile = tile[..., :-latent_y_overlap[i], :] + if j < len(row) - 1: + tile = tile[..., :, :-latent_x_overlap[j]] + result_row.append(tile) + result_rows.append(torch.cat(result_row, dim=-1)) + return torch.cat(result_rows, dim=-2) + + def tiled_decode(self, z): + height, width = z.shape[-2] * self.vae_ratio, z.shape[-1] * self.vae_ratio + y_idx, y_len, y_overlap = self.split_tiles(height) + x_idx, x_len, x_overlap = self.split_tiles(width) + + # Blended tiles are written straight into a pre-allocated canvas. + canvas = None + row_tails = [] + out_y = 0 + for i, (i_pos, i_len) in enumerate(zip(y_idx, y_len)): + zi, zl = i_pos // self.vae_ratio, i_len // self.vae_ratio + new_tails = [] + left_tail = None + out_x = 0 + for j, (j_pos, j_len) in enumerate(zip(x_idx, x_len)): + zj, zw = j_pos // self.vae_ratio, j_len // self.vae_ratio + tile = self._decode_pixels(z[..., zi:zi + zl, zj:zj + zw]) + if i < len(y_idx) - 1: + new_tails.append(tile[..., -y_overlap[i]:, :].clone()) + next_left_tail = tile[..., :, -x_overlap[j]:].clone() if j < len(x_idx) - 1 else None + if i > 0: + tile = self.blend(row_tails[j], tile, y_overlap[i - 1], dim=-2) + if j > 0: + tile = self.blend(left_tail, tile, x_overlap[j - 1], dim=-1) + left_tail = next_left_tail + if i < len(y_idx) - 1: + tile = tile[..., :-y_overlap[i], :] + if j < len(x_idx) - 1: + tile = tile[..., :, :-x_overlap[j]] + if canvas is None: + canvas = torch.empty(*tile.shape[:-2], height, width, dtype=tile.dtype, device=tile.device) + canvas[..., out_y:out_y + tile.shape[-2], out_x:out_x + tile.shape[-1]].copy_(tile) + out_x += tile.shape[-1] + row_tails = new_tails + out_y += tile.shape[-2] + return canvas + + # 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 + + 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)) + + z = torch.cat(z_list, dim=2) + if self.token_drop > 0: + z = z[:, :, :-self.token_drop] + return z + + def _decode_temporal_pad_frames(self, z_len, pad_tokens): + if pad_tokens <= 0: + return 0 + intra_tail = self.clip_length % self.vae_ratio_t + if intra_tail == 0: + return pad_tokens * self.vae_ratio_t + + z_len_before_pad = z_len - pad_tokens + return sum( + (intra_tail if (z_len_before_pad + k) % self.tokens_chunk_size == 0 + else self.vae_ratio_t) + for k in range(pad_tokens) + ) + + def _decode_temporal_frame_plan(self, z_len, num_chunks, pad_tokens): + chunk_dec = self.tokens_chunk_size * self.vae_ratio_t + split_count = int(self.token_drop > 0) + 1 + total_frames = 0 + final_overlap_frames = 0 + + for i in range(num_chunks): + t_start_idx = i * self.tokens_chunk_size + t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap + clip_token_len = max(0, min(t_end_idx, z_len) - min(t_start_idx, z_len)) + clip_frame_len = clip_token_len * self.vae_ratio_t + + for j in range(split_count): + f_start_idx = j * chunk_dec + f_end_idx = min(f_start_idx + chunk_dec, clip_frame_len) + chunk_frames = max(0, f_end_idx - f_start_idx - self.frame_pre_padding) + if j == 0: + total_frames += chunk_frames + else: + final_overlap_frames = chunk_frames + + 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 + + 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 + + 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_overlap = None + write_pos = 0 + + def write_part(part): + nonlocal dec, 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) + copy_frames = min(part_frames, max(0, dec.shape[2] - write_pos)) + if copy_frames > 0: + dec[:, :, write_pos:write_pos + copy_frames, :, :].copy_( + part[:, :, :copy_frames, :, :] + ) + write_pos += copy_frames + + for i in range(num_chunks): + t_start_idx = i * self.tokens_chunk_size + t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap + clip_z = z[:, :, t_start_idx:t_end_idx, :, :] + + clip_dec = self._adaptive_decode(clip_z) + + for j in range(split_count): + f_start_idx = j * chunk_dec + f_end_idx = min(f_start_idx + chunk_dec, clip_dec.shape[2]) + clip_dec_chunk = clip_dec[:, :, f_start_idx:f_end_idx, :, :] + clip_dec_chunk = clip_dec_chunk[:, :, self.frame_pre_padding:, :, :] + + if j == 0: + if dec_overlap is not None: + clip_dec_chunk = self.blend( + dec_overlap, clip_dec_chunk, self.frame_overlap, dim=-3 + ) + dec_overlap = None + write_part(clip_dec_chunk) + else: + dec_overlap = clip_dec_chunk.contiguous() + + if i == num_chunks - 1 and dec_overlap is not None: + write_part(dec_overlap) + dec_overlap = None + + del clip_dec, clip_z + + return dec + + + def encode(self, x): + # 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 x.shape[2] == 1: + moments = self._adaptive_encode(x) + moments = moments[:, :, -1:, :, :] + else: + moments = self.encode_temporal(x) + + mean = torch.chunk(moments.float(), 2, dim=1)[0] + + latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(mean) + latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(mean) + return (mean - latents_mean) / latents_std + + def encode_tiled(self, x, **kwargs): + # tiling is always on internally with the reference's semantic tile sizes, ignore tiling fallbacks + return self.encode(x) + + 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] + 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 diff --git a/comfy/ldm/modules/attention.py b/comfy/ldm/modules/attention.py index 55360535a..2c549e095 100644 --- a/comfy/ldm/modules/attention.py +++ b/comfy/ldm/modules/attention.py @@ -1,5 +1,6 @@ import math import sys +import inspect import torch import torch.nn.functional as F @@ -14,16 +15,16 @@ from .sub_quadratic_attention import efficient_dot_product_attention from comfy import model_management -TORCH_HAS_GQA = model_management.torch_version_numeric >= (2, 5) - if model_management.xformers_enabled(): import xformers import xformers.ops SAGE_ATTENTION_IS_AVAILABLE = False +SAGE_ATTENTION_SUPPORTS_MASK = False try: from sageattention import sageattn SAGE_ATTENTION_IS_AVAILABLE = True + SAGE_ATTENTION_SUPPORTS_MASK = "attn_mask" in inspect.signature(sageattn).parameters except ImportError as e: if model_management.sage_attention_enabled(): if e.name == "sageattention": @@ -89,6 +90,26 @@ def default(val, d): return val return d +def _heads_from_dim(tensor, dim_head, name): + inner_dim = tensor.shape[-1] + if inner_dim % dim_head != 0: + raise ValueError(f"{name} inner dimension {inner_dim} is not divisible by head dimension {dim_head}") + return inner_dim // dim_head + +def _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, enable_gqa=False, expand_kv=True): + q = q.unsqueeze(3).reshape(b, -1, heads, dim_head) + if enable_gqa: + key_heads = _heads_from_dim(k, dim_head, "Key") + value_heads = _heads_from_dim(v, dim_head, "Value") + else: + key_heads = heads + value_heads = heads + k = k.unsqueeze(3).reshape(b, -1, key_heads, dim_head) + v = v.unsqueeze(3).reshape(b, -1, value_heads, dim_head) + if enable_gqa and expand_kv: + k, v = comfy.ops.repeat_kv_for_gqa(k, v, heads, -2) + return q, k, v + # feedforward class GEGLU(nn.Module): @@ -152,28 +173,19 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape b, _, dim_head = q.shape dim_head //= heads - if kwargs.get("enable_gqa", False) and q.shape[-3] != k.shape[-3]: - n_rep = q.shape[-3] // k.shape[-3] - k = k.repeat_interleave(n_rep, dim=-3) - v = v.repeat_interleave(n_rep, dim=-3) - scale = kwargs.get("scale", dim_head ** -0.5) h = heads if skip_reshape: - q, k, v = map( + if kwargs.get("enable_gqa", False): + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) + q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), ) else: - q, k, v = map( - lambda t: t.unsqueeze(3) - .reshape(b, -1, heads, dim_head) - .permute(0, 2, 1, 3) - .reshape(b * heads, -1, dim_head) - .contiguous(), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) + q, k, v = map(lambda t: t.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(), (q, k, v)) # force cast to fp32 to avoid overflowing if attn_precision == torch.float32: @@ -231,13 +243,16 @@ def attention_sub_quad(query, key, value, heads, mask=None, attn_precision=None, query = query * (kwargs["scale"] * dim_head ** 0.5) if skip_reshape: + if kwargs.get("enable_gqa", False): + key, value = comfy.ops.repeat_kv_for_gqa(key, value, query.shape[-3], -3) query = query.reshape(b * heads, -1, dim_head) value = value.reshape(b * heads, -1, dim_head) key = key.reshape(b * heads, -1, dim_head).movedim(1, 2) else: - query = query.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) - value = value.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) - key = key.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 3, 1).reshape(b * heads, dim_head, -1) + query, key, value = _reshape_qkv_to_heads(query, key, value, b, heads, dim_head, kwargs.get("enable_gqa", False)) + query = query.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) + value = value.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head) + key = key.permute(0, 2, 3, 1).reshape(b * heads, dim_head, -1) dtype = query.dtype @@ -304,19 +319,15 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape scale = kwargs.get("scale", dim_head ** -0.5) if skip_reshape: - q, k, v = map( + if kwargs.get("enable_gqa", False): + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) + q, k, v = map( lambda t: t.reshape(b * heads, -1, dim_head), (q, k, v), ) else: - q, k, v = map( - lambda t: t.unsqueeze(3) - .reshape(b, -1, heads, dim_head) - .permute(0, 2, 1, 3) - .reshape(b * heads, -1, dim_head) - .contiguous(), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) + q, k, v = map(lambda t: t.permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(), (q, k, v)) r1 = torch.zeros(q.shape[0], q.shape[1], v.shape[2], device=q.device, dtype=q.dtype) @@ -438,7 +449,7 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh disabled_xformers = True if disabled_xformers: - return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape, **kwargs) + return attention_pytorch(q, k, v, heads, mask, skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs) if skip_reshape: # b h k d -> b k h d @@ -446,13 +457,12 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh lambda t: t.permute(0, 2, 1, 3), (q, k, v), ) + if kwargs.get("enable_gqa", False): + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-2], -2) # actually do the reshaping else: dim_head //= heads - q, k, v = map( - lambda t: t.reshape(b, -1, heads, dim_head), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) if mask is not None: # add a singleton batch dimension @@ -474,7 +484,7 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh mask = mask_out[..., :mask.shape[-1]] mask = mask.expand(b, heads, -1, -1) - out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask) + out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask, scale=kwargs.get("scale", None)) if skip_output_reshape: out = out.permute(0, 2, 1, 3) @@ -498,10 +508,8 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False), expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) if mask is not None: # add a batch dimension if there isn't already one @@ -511,9 +519,7 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha if mask.ndim == 3: mask = mask.unsqueeze(1) - # Pass through extra SDPA kwargs (scale, enable_gqa) if provided - # enable_gqa requires PyTorch 2.5+; older versions use manual KV expansion above - sdpa_keys = ("scale", "enable_gqa") if TORCH_HAS_GQA else ("scale",) + sdpa_keys = ("scale", "enable_gqa") sdpa_extra = {k: v for k, v in kwargs.items() if k in sdpa_keys} if SDP_BATCH_LIMIT >= b: @@ -541,20 +547,19 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha @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: + if kwargs.get("low_precision_attention", True) is False or (mask is not None and not SAGE_ATTENTION_SUPPORTS_MASK): return attention_pytorch(q, k, v, heads, mask=mask, skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs) exception_fallback = False if skip_reshape: b, _, _, dim_head = q.shape tensor_layout = "HND" + if kwargs.get("enable_gqa", False): + k, v = comfy.ops.repeat_kv_for_gqa(k, v, q.shape[-3], -3) else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False)) tensor_layout = "NHD" if mask is not None: @@ -565,8 +570,12 @@ def attention_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape= if mask.ndim == 3: mask = mask.unsqueeze(1) + sage_kwargs = {"is_causal": False, "tensor_layout": tensor_layout, "sm_scale": kwargs.get("scale", None), "smooth_k": False} + if mask is not None: + sage_kwargs["attn_mask"] = mask + try: - out = sageattn(q, k, v, attn_mask=mask, is_causal=False, tensor_layout=tensor_layout) + out = sageattn(q, k, v, **sage_kwargs) except Exception as e: logging.error("Error running sage attention: {}, using pytorch attention instead.".format(e)) exception_fallback = True @@ -616,7 +625,6 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape skip_output_reshape=skip_output_reshape, **kwargs ) - q_s, k_s, v_s = q, k, v N = q.shape[2] dim_head = D else: @@ -642,11 +650,15 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape **kwargs ) - if not skip_reshape: - q_s, k_s, v_s = map( - lambda t: t.view(B, -1, heads, dim_head).permute(0, 2, 1, 3).contiguous(), - (q, k, v), - ) + if skip_reshape: + q_s = q + if kwargs.get("enable_gqa", False): + k_s, v_s = comfy.ops.repeat_kv_for_gqa(k, v, H, -3) + else: + k_s, v_s = k, v + else: + q_s, k_s, v_s = _reshape_qkv_to_heads(q, k, v, B, heads, dim_head, kwargs.get("enable_gqa", False)) + q_s, k_s, v_s = map(lambda t: t.permute(0, 2, 1, 3).contiguous(), (q_s, k_s, v_s)) B, H, L, D = q_s.shape try: @@ -662,7 +674,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape q, k, v, heads, mask=mask, attn_precision=attn_precision, - skip_reshape=False, + skip_reshape=skip_reshape, skip_output_reshape=skip_output_reshape, **kwargs ) @@ -679,21 +691,22 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape return out try: - @torch.library.custom_op("flash_attention::flash_attn", mutates_args=()) + @torch.library.custom_op("comfy::flash_attn", mutates_args=()) def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - dropout_p: float = 0.0, causal: bool = False) -> torch.Tensor: - return flash_attn_func(q, k, v, dropout_p=dropout_p, causal=causal) + dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor: + softmax_scale_arg = None if softmax_scale == -1.0 else softmax_scale + return flash_attn_func(q, k, v, dropout_p=dropout_p, causal=causal, softmax_scale=softmax_scale_arg) @flash_attn_wrapper.register_fake - def flash_attn_fake(q, k, v, dropout_p=0.0, causal=False): + def flash_attn_fake(q, k, v, dropout_p=0.0, causal=False, softmax_scale=-1.0): # Output shape is the same as q return q.new_empty(q.shape) except AttributeError as error: FLASH_ATTN_ERROR = error def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - dropout_p: float = 0.0, causal: bool = False) -> torch.Tensor: + dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor: assert False, f"Could not define flash_attn_wrapper: {FLASH_ATTN_ERROR}" @wrap_attn @@ -703,10 +716,8 @@ def attention_flash(q, k, v, heads, mask=None, attn_precision=None, skip_reshape else: b, _, dim_head = q.shape dim_head //= heads - q, k, v = map( - lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2), - (q, k, v), - ) + q, k, v = _reshape_qkv_to_heads(q, k, v, b, heads, dim_head, kwargs.get("enable_gqa", False), expand_kv=False) + q, k, v = map(lambda t: t.transpose(1, 2), (q, k, v)) if mask is not None: # add a batch dimension if there isn't already one @@ -725,10 +736,16 @@ def attention_flash(q, k, v, heads, mask=None, attn_precision=None, skip_reshape v.transpose(1, 2), dropout_p=0.0, causal=False, + softmax_scale=kwargs.get("scale", -1.0), ).transpose(1, 2) except Exception as e: logging.warning(f"Flash Attention failed, using default SDPA: {e}") - out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False) + sdpa_extra = {} + if kwargs.get("enable_gqa", False): + sdpa_extra["enable_gqa"] = True + if "scale" in kwargs: + sdpa_extra["scale"] = kwargs["scale"] + out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False, **sdpa_extra) if not skip_output_reshape: out = ( out.transpose(1, 2).reshape(b, -1, heads * dim_head) @@ -1209,5 +1226,3 @@ class SpatialVideoTransformer(SpatialTransformer): x = self.proj_out(x) out = x + x_in return out - - diff --git a/comfy/ldm/modules/diffusionmodules/model.py b/comfy/ldm/modules/diffusionmodules/model.py index fcbaa074f..e752d0ecb 100644 --- a/comfy/ldm/modules/diffusionmodules/model.py +++ b/comfy/ldm/modules/diffusionmodules/model.py @@ -22,7 +22,7 @@ def torch_cat_if_needed(xl, dim): else: return None -def get_timestep_embedding(timesteps, embedding_dim): +def get_timestep_embedding(timesteps, embedding_dim, flip_sin_to_cos=False, downscale_freq_shift=1): """ This matches the implementation in Denoising Diffusion Probabilistic Models: From Fairseq. @@ -33,11 +33,13 @@ def get_timestep_embedding(timesteps, embedding_dim): assert len(timesteps.shape) == 1 half_dim = embedding_dim // 2 - emb = math.log(10000) / (half_dim - 1) + emb = math.log(10000) / (half_dim - downscale_freq_shift) emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb) emb = emb.to(device=timesteps.device) emb = timesteps.float()[:, None] * emb[None, :] emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) if embedding_dim % 2 == 1: # zero pad emb = torch.nn.functional.pad(emb, (0,1,0,0)) return emb diff --git a/comfy/ldm/omnigen/omnigen2.py b/comfy/ldm/omnigen/omnigen2.py index b8da4cf39..d18a9f6d0 100644 --- a/comfy/ldm/omnigen/omnigen2.py +++ b/comfy/ldm/omnigen/omnigen2.py @@ -141,11 +141,8 @@ class Attention(nn.Module): key = key.transpose(1, 2) value = value.transpose(1, 2) - if self.kv_heads < self.heads: - key = key.repeat_interleave(self.heads // self.kv_heads, dim=1) - value = value.repeat_interleave(self.heads // self.kv_heads, dim=1) - - hidden_states = optimized_attention_masked(query, key, value, self.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options) + gqa_kwargs = {"enable_gqa": True} if self.kv_heads < self.heads else {} + hidden_states = optimized_attention_masked(query, key, value, self.heads, attention_mask, skip_reshape=True, transformer_options=transformer_options, **gqa_kwargs) hidden_states = self.to_out[0](hidden_states) return hidden_states diff --git a/comfy/ldm/pixeldit/model.py b/comfy/ldm/pixeldit/model.py index b044b9b29..3b30b9226 100644 --- a/comfy/ldm/pixeldit/model.py +++ b/comfy/ldm/pixeldit/model.py @@ -197,6 +197,9 @@ class PixDiT_T2I(nn.Module): """Hook for subclasses to inject per-block state into the patch stream (e.g. PiD's LQ gate).""" return s + def _pre_pixel_blocks(self, s, **kwargs): + return s + def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, **kwargs): H_orig, W_orig = x.shape[2], x.shape[3] x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size)) @@ -226,6 +229,7 @@ class PixDiT_T2I(nn.Module): s, y_emb = blk(s, y_emb, condition, pos_img, pos_txt, None, transformer_options=transformer_options) s = F.silu(t_emb + s) + s = self._pre_pixel_blocks(s, **kwargs) s_cond = s.view(B * L, self.hidden_size) x_pixels = self.pixel_embedder(x, patch_size=self.patch_size) for blk in self.pixel_blocks: diff --git a/comfy/ldm/pixeldit/pid.py b/comfy/ldm/pixeldit/pid.py index 21b73907a..8590408d9 100644 --- a/comfy/ldm/pixeldit/pid.py +++ b/comfy/ldm/pixeldit/pid.py @@ -13,15 +13,15 @@ from .model import PixDiT_T2I from .modules import precompute_freqs_cis_2d -class SigmaAwareGatePerTokenPerDim(nn.Module): +class SigmaAwareGate(nn.Module): """gate = sigmoid(content_proj(cat[x, lq]) - exp(log_alpha) * sigma); out = x + gate * lq. Trained init gives ~0.88 gate at sigma=0, ~0.05 at sigma=1. """ - def __init__(self, dim: int, dtype=None, device=None, operations=None): + def __init__(self, dim: int, per_token: bool = False, dtype=None, device=None, operations=None): super().__init__() - self.content_proj = operations.Linear(dim * 2, dim, dtype=dtype, device=device) + self.content_proj = operations.Linear(dim * 2, 1 if per_token else dim, dtype=dtype, device=device) self.log_alpha = nn.Parameter(torch.empty((), dtype=dtype, device=device)) def forward(self, x: torch.Tensor, lq: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: @@ -36,15 +36,15 @@ class SigmaAwareGatePerTokenPerDim(nn.Module): class ResBlock(nn.Module): """Pre-activation ResNet block: GN -> SiLU -> Conv -> GN -> SiLU -> Conv + skip.""" - def __init__(self, channels: int, num_groups: int = 4, dtype=None, device=None, operations=None): + def __init__(self, channels: int, num_groups: int = 4, conv_padding_mode: str = "zeros", dtype=None, device=None, operations=None): super().__init__() self.block = nn.Sequential( operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), ) def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -62,9 +62,13 @@ class LQProjection2D(nn.Module): patch_size: int = 16, sr_scale: int = 4, latent_spatial_down_factor: int = 8, + latent_unpatchify_factor: int = 1, num_res_blocks: int = 4, num_outputs: int = 7, interval: int = 2, + conv_padding_mode: str = "zeros", + gate_per_token: bool = False, + pit_output: bool = False, dtype=None, device=None, operations=None, ): super().__init__() @@ -74,34 +78,38 @@ class LQProjection2D(nn.Module): self.patch_size = patch_size self.sr_scale = sr_scale self.latent_spatial_down_factor = latent_spatial_down_factor + self.latent_unpatchify_factor = latent_unpatchify_factor self.num_outputs = num_outputs self.interval = interval - z_to_patch_ratio = (sr_scale * latent_spatial_down_factor) / patch_size + effective_latent_channels = latent_channels // (latent_unpatchify_factor * latent_unpatchify_factor) + effective_spatial_down_factor = latent_spatial_down_factor // latent_unpatchify_factor + z_to_patch_ratio = (sr_scale * effective_spatial_down_factor) / patch_size self.z_to_patch_ratio = z_to_patch_ratio if z_to_patch_ratio >= 1: self.latent_fold_factor = 0 - latent_proj_in_ch = latent_channels + latent_proj_in_ch = effective_latent_channels else: fold_factor = int(1 / z_to_patch_ratio) assert fold_factor * z_to_patch_ratio == 1.0 self.latent_fold_factor = fold_factor - latent_proj_in_ch = latent_channels * fold_factor * fold_factor + latent_proj_in_ch = effective_latent_channels * fold_factor * fold_factor layers = [ - operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), nn.SiLU(), - operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device), ] for _ in range(num_res_blocks): - layers.append(ResBlock(hidden_dim, dtype=dtype, device=device, operations=operations)) + layers.append(ResBlock(hidden_dim, conv_padding_mode=conv_padding_mode, dtype=dtype, device=device, operations=operations)) self.latent_proj = nn.Sequential(*layers) self.output_heads = nn.ModuleList( [operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) for _ in range(num_outputs)] ) + self.pit_head = operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) if pit_output else None self.gate_modules = nn.ModuleList( - [SigmaAwareGatePerTokenPerDim(out_dim, dtype=dtype, device=device, operations=operations) + [SigmaAwareGate(out_dim, per_token=gate_per_token, dtype=dtype, device=device, operations=operations) for _ in range(num_outputs)] ) @@ -115,6 +123,11 @@ class LQProjection2D(nn.Module): return self.gate_modules[out_idx](x, lq_feature, sigma) def _align_latent_to_patch_grid(self, lq_latent: torch.Tensor, pH: int, pW: int) -> torch.Tensor: + f = self.latent_unpatchify_factor + if f > 1: + B, C, H, W = lq_latent.shape + lq_latent = lq_latent.reshape(B, C // (f * f), f, f, H, W) + lq_latent = lq_latent.permute(0, 1, 4, 2, 5, 3).reshape(B, C // (f * f), H * f, W * f) B, z_dim = lq_latent.shape[:2] if self.z_to_patch_ratio >= 1: if lq_latent.shape[2] != pH or lq_latent.shape[3] != pW: @@ -134,7 +147,10 @@ class LQProjection2D(nn.Module): feat = self._align_latent_to_patch_grid(lq_latent, target_pH, target_pW) B, C, H, W = feat.shape tokens = feat.permute(0, 2, 3, 1).contiguous().view(B, H * W, C) - return [head(tokens) for head in self.output_heads] + outputs = [head(tokens) for head in self.output_heads] + if self.pit_head is not None: + outputs.append(self.pit_head(tokens)) + return outputs class PidNet(PixDiT_T2I): @@ -148,6 +164,10 @@ class PidNet(PixDiT_T2I): lq_interval: int = 2, sr_scale: int = 4, latent_spatial_down_factor: int = 8, + lq_latent_unpatchify_factor: int = 1, + lq_conv_padding_mode: str = "zeros", + lq_gate_per_token: bool = False, + pit_lq_inject: bool = False, rope_ref_h: int = 1024, # NTK ref resolution in PIXEL units: 1024px / patch=16 -> grid_ref=64. rope_ref_w: int = 1024, image_model=None, @@ -165,6 +185,8 @@ class PidNet(PixDiT_T2I): for blk in self.pixel_blocks: blk._rope_fn = _pit_rope_fn + self.pit_lq_inject = pit_lq_inject + num_lq_outputs = (self.patch_depth + lq_interval - 1) // lq_interval self.lq_proj = LQProjection2D( latent_channels=lq_latent_channels, @@ -173,13 +195,20 @@ class PidNet(PixDiT_T2I): patch_size=self.patch_size, sr_scale=sr_scale, latent_spatial_down_factor=latent_spatial_down_factor, + latent_unpatchify_factor=lq_latent_unpatchify_factor, num_res_blocks=lq_num_res_blocks, num_outputs=num_lq_outputs, interval=lq_interval, + conv_padding_mode=lq_conv_padding_mode, + gate_per_token=lq_gate_per_token, + pit_output=pit_lq_inject, dtype=dtype, device=device, operations=operations, ) + self.pit_lq_gate = SigmaAwareGate( + self.hidden_size, per_token=lq_gate_per_token, dtype=dtype, device=device, operations=operations + ) if pit_lq_inject else None def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts): return precompute_freqs_cis_2d( @@ -197,6 +226,11 @@ class PidNet(PixDiT_T2I): return s return self.lq_proj.gate(s, pid_lq_features[out_idx], pid_degrade_sigma, out_idx) + def _pre_pixel_blocks(self, s, pid_pit_lq_feature=None, pid_degrade_sigma=None, **kwargs): + if pid_pit_lq_feature is None: + return s + return self.pit_lq_gate(s, pid_pit_lq_feature, pid_degrade_sigma) + def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, lq_latent=None, degrade_sigma=None, **kwargs): if lq_latent is None: raise ValueError("PidNet requires lq_latent — attach via PiDConditioning") @@ -216,12 +250,14 @@ class PidNet(PixDiT_T2I): degrade_sigma = degrade_sigma.expand(B).contiguous() lq_features = self.lq_proj(lq_latent=lq_latent.to(x), target_pH=Hs, target_pW=Ws) + pit_lq_feature = lq_features.pop() if self.pit_lq_inject else None return super()._forward( x, timesteps, context=context, attention_mask=attention_mask, transformer_options=transformer_options, pid_lq_features=lq_features, + pid_pit_lq_feature=pit_lq_feature, pid_degrade_sigma=degrade_sigma, **kwargs, ) diff --git a/comfy/ldm/seedvr/attention.py b/comfy/ldm/seedvr/attention.py new file mode 100644 index 000000000..11b4c1e4a --- /dev/null +++ b/comfy/ldm/seedvr/attention.py @@ -0,0 +1,51 @@ +import torch + +from comfy.ldm.modules import attention as _attention + + +def _var_attention_qkv(q, k, v, heads, skip_reshape): + if skip_reshape: + return q, k, v, q.shape[-1] + total_tokens, embed_dim = q.shape + head_dim = embed_dim // heads + return ( + q.view(total_tokens, heads, head_dim), + k.view(k.shape[0], heads, head_dim), + v.view(v.shape[0], heads, head_dim), + head_dim, + ) + + +def _var_attention_output(out, heads, head_dim, skip_output_reshape): + if skip_output_reshape: + return out + return out.reshape(-1, heads * head_dim) + + +def var_attention_optimized_split(q, k, v, heads, cu_seqlens_q, cu_seqlens_k, *args, skip_reshape=False, skip_output_reshape=False, **kwargs): + q, k, v, head_dim = _var_attention_qkv(q, k, v, heads, skip_reshape) + + q_split_indices = cu_seqlens_q[1:-1] + k_split_indices = cu_seqlens_k[1:-1] + if k.shape[0] != v.shape[0]: + raise ValueError("cu_seqlens_k does not match v token count") + + q_splits = torch.tensor_split(q, q_split_indices, dim=0) + k_splits = torch.tensor_split(k, k_split_indices, dim=0) + v_splits = torch.tensor_split(v, k_split_indices, dim=0) + if len(q_splits) != len(k_splits) or len(q_splits) != len(v_splits): + raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same sequence count") + + out = [] + for q_i, k_i, v_i in zip(q_splits, k_splits, v_splits): + q_i = q_i.permute(1, 0, 2).unsqueeze(0) + k_i = k_i.permute(1, 0, 2).unsqueeze(0) + v_i = v_i.permute(1, 0, 2).unsqueeze(0) + out_i = _attention.optimized_attention(q_i, k_i, v_i, heads, skip_reshape=True, skip_output_reshape=True) + out.append(out_i.squeeze(0).permute(1, 0, 2)) + + out = torch.cat(out, dim=0) + return _var_attention_output(out, heads, head_dim, skip_output_reshape) + + +optimized_var_attention = var_attention_optimized_split diff --git a/comfy/ldm/seedvr/color_fix.py b/comfy/ldm/seedvr/color_fix.py new file mode 100644 index 000000000..a43cb5270 --- /dev/null +++ b/comfy/ldm/seedvr/color_fix.py @@ -0,0 +1,301 @@ +import torch +import torch.nn.functional as F +from torch import Tensor + +from comfy.ldm.seedvr.constants import ( + CIELAB_DELTA, + CIELAB_KAPPA, + D65_WHITE_X, + D65_WHITE_Z, + WAVELET_DECOMP_LEVELS, +) + + +def wavelet_blur(image: Tensor, radius): + max_safe_radius = max(1, min(image.shape[-2:]) // 8) + if radius > max_safe_radius: + radius = max_safe_radius + + num_channels = image.shape[1] + + kernel_vals = [ + [0.0625, 0.125, 0.0625], + [0.125, 0.25, 0.125], + [0.0625, 0.125, 0.0625], + ] + kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device) + kernel = kernel[None, None].repeat(num_channels, 1, 1, 1) + + image = F.pad(image, (radius, radius, radius, radius), mode='replicate') + output = F.conv2d(image, kernel, groups=num_channels, dilation=radius) + + return output + +def wavelet_decomposition(image: Tensor, levels: int = WAVELET_DECOMP_LEVELS): + high_freq = torch.zeros_like(image) + + for i in range(levels): + radius = 2 ** i + low_freq = wavelet_blur(image, radius) + high_freq.add_(image).sub_(low_freq) + image = low_freq + + return high_freq, low_freq + +def wavelet_reconstruction(content_feat: Tensor, style_feat: Tensor) -> Tensor: + + if content_feat.shape != style_feat.shape: + if len(content_feat.shape) >= 3: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False + ) + + content_high_freq, content_low_freq = wavelet_decomposition(content_feat) + del content_low_freq + + style_high_freq, style_low_freq = wavelet_decomposition(style_feat) + del style_high_freq + + if content_high_freq.shape != style_low_freq.shape: + style_low_freq = F.interpolate( + style_low_freq, + size=content_high_freq.shape[-2:], + mode='bilinear', + align_corners=False + ) + + content_high_freq.add_(style_low_freq) + + return content_high_freq.clamp_(-1.0, 1.0) + +def _histogram_matching_channel(source: Tensor, reference: Tensor) -> Tensor: + original_shape = source.shape + + source_flat = source.flatten() + reference_flat = reference.flatten() + + source_sorted, source_indices = torch.sort(source_flat) + reference_sorted, _ = torch.sort(reference_flat) + del reference_flat + + n_source = len(source_sorted) + n_reference = len(reference_sorted) + + if n_source == n_reference: + matched_sorted = reference_sorted + else: + source_quantiles = torch.linspace(0, 1, n_source, device=source.device) + ref_indices = (source_quantiles * (n_reference - 1)).long() + ref_indices.clamp_(0, n_reference - 1) + matched_sorted = reference_sorted[ref_indices] + del source_quantiles, ref_indices, reference_sorted + + del source_sorted, source_flat + + inverse_indices = torch.argsort(source_indices) + del source_indices + matched_flat = matched_sorted[inverse_indices] + del matched_sorted, inverse_indices + + return matched_flat.reshape(original_shape) + +def _lab_to_rgb_batch(lab: Tensor, matrix_inv: Tensor, epsilon: float, kappa: float) -> Tensor: + L, a, b = lab[:, 0], lab[:, 1], lab[:, 2] + + fy = (L + 16.0) / 116.0 + fx = a.div(500.0).add_(fy) + fz = fy - b / 200.0 + del L, a, b + + x = torch.where( + fx > epsilon, + torch.pow(fx, 3.0), + fx.mul(116.0).sub_(16.0).div_(kappa) + ) + y = torch.where( + fy > epsilon, + torch.pow(fy, 3.0), + fy.mul(116.0).sub_(16.0).div_(kappa) + ) + z = torch.where( + fz > epsilon, + torch.pow(fz, 3.0), + fz.mul(116.0).sub_(16.0).div_(kappa) + ) + del fx, fy, fz + + x.mul_(D65_WHITE_X) + z.mul_(D65_WHITE_Z) + + xyz = torch.stack([x, y, z], dim=1) + del x, y, z + + B, _, H, W = xyz.shape + xyz_flat = xyz.permute(0, 2, 3, 1).reshape(-1, 3) + del xyz + + xyz_flat = xyz_flat.to(dtype=matrix_inv.dtype) + rgb_linear_flat = torch.matmul(xyz_flat, matrix_inv.T) + del xyz_flat + + rgb_linear = rgb_linear_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2) + del rgb_linear_flat + + mask = rgb_linear > 0.0031308 + rgb = torch.where( + mask, + torch.pow(torch.clamp(rgb_linear, min=0.0), 1.0 / 2.4).mul_(1.055).sub_(0.055), + rgb_linear * 12.92 + ) + del mask, rgb_linear + + return torch.clamp(rgb, 0.0, 1.0) + +def _rgb_to_lab_batch(rgb: Tensor, matrix: Tensor, epsilon: float, kappa: float) -> Tensor: + mask = rgb > 0.04045 + rgb_linear = torch.where( + mask, + torch.pow((rgb + 0.055) / 1.055, 2.4), + rgb / 12.92 + ) + del mask + + B, _, H, W = rgb_linear.shape + rgb_flat = rgb_linear.permute(0, 2, 3, 1).reshape(-1, 3) + del rgb_linear + + rgb_flat = rgb_flat.to(dtype=matrix.dtype) + xyz_flat = torch.matmul(rgb_flat, matrix.T) + del rgb_flat + + xyz = xyz_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2) + del xyz_flat + + xyz[:, 0].div_(D65_WHITE_X) + xyz[:, 2].div_(D65_WHITE_Z) + + epsilon_cubed = epsilon ** 3 + mask = xyz > epsilon_cubed + f_xyz = torch.where( + mask, + torch.pow(xyz, 1.0 / 3.0), + xyz.mul(kappa).add_(16.0).div_(116.0) + ) + del xyz, mask + + L = f_xyz[:, 1].mul(116.0).sub_(16.0) + a = (f_xyz[:, 0] - f_xyz[:, 1]).mul_(500.0) + b = (f_xyz[:, 1] - f_xyz[:, 2]).mul_(200.0) + del f_xyz + + return torch.stack([L, a, b], dim=1) + +def lab_color_transfer( + content_feat: Tensor, + style_feat: Tensor, + luminance_weight: float = 0.8 +) -> Tensor: + content_feat = wavelet_reconstruction(content_feat, style_feat) + + if content_feat.shape != style_feat.shape: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False + ) + + device = content_feat.device + original_dtype = content_feat.dtype + content_feat = content_feat.float() + style_feat = style_feat.float() + + rgb_to_xyz_matrix = torch.tensor([ + [0.4124564, 0.3575761, 0.1804375], + [0.2126729, 0.7151522, 0.0721750], + [0.0193339, 0.1191920, 0.9503041] + ], dtype=torch.float32, device=device) + + xyz_to_rgb_matrix = torch.tensor([ + [ 3.2404542, -1.5371385, -0.4985314], + [-0.9692660, 1.8760108, 0.0415560], + [ 0.0556434, -0.2040259, 1.0572252] + ], dtype=torch.float32, device=device) + + epsilon = CIELAB_DELTA + kappa = CIELAB_KAPPA + + content_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) + style_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) + + content_lab = _rgb_to_lab_batch(content_feat, rgb_to_xyz_matrix, epsilon, kappa) + del content_feat + + style_lab = _rgb_to_lab_batch(style_feat, rgb_to_xyz_matrix, epsilon, kappa) + del style_feat, rgb_to_xyz_matrix + + matched_a = _histogram_matching_channel(content_lab[:, 1], style_lab[:, 1]) + matched_b = _histogram_matching_channel(content_lab[:, 2], style_lab[:, 2]) + + if luminance_weight < 1.0: + matched_L = _histogram_matching_channel(content_lab[:, 0], style_lab[:, 0]) + result_L = content_lab[:, 0].mul(luminance_weight).add_(matched_L.mul(1.0 - luminance_weight)) + del matched_L + else: + result_L = content_lab[:, 0] + + del content_lab, style_lab + + result_lab = torch.stack([result_L, matched_a, matched_b], dim=1) + del result_L, matched_a, matched_b + + result_rgb = _lab_to_rgb_batch(result_lab, xyz_to_rgb_matrix, epsilon, kappa) + del result_lab, xyz_to_rgb_matrix + + result = result_rgb.mul_(2.0).sub_(1.0) + del result_rgb + + result = result.to(original_dtype) + + return result + + +def wavelet_color_transfer(content_feat: Tensor, style_feat: Tensor) -> Tensor: + return wavelet_reconstruction(content_feat, style_feat) + + +def adain_color_transfer(content_feat: Tensor, style_feat: Tensor, eps: float = 1e-5) -> Tensor: + if content_feat.shape != style_feat.shape: + style_feat = F.interpolate( + style_feat, + size=content_feat.shape[-2:], + mode='bilinear', + align_corners=False, + ) + + original_dtype = content_feat.dtype + content_feat = content_feat.float() + style_feat = style_feat.float() + + b, c = content_feat.shape[:2] + content_flat = content_feat.reshape(b, c, -1) + style_flat = style_feat.reshape(b, c, -1) + + content_mean = content_flat.mean(dim=2).reshape(b, c, 1, 1) + content_std = (content_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1) + style_mean = style_flat.mean(dim=2).reshape(b, c, 1, 1) + style_std = (style_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1) + del content_flat, style_flat + + normalized = (content_feat - content_mean) / content_std + del content_mean, content_std + result = normalized * style_std + style_mean + del normalized, style_mean, style_std + + result = result.clamp_(-1.0, 1.0) + if result.dtype != original_dtype: + result = result.to(original_dtype) + return result diff --git a/comfy/ldm/seedvr/constants.py b/comfy/ldm/seedvr/constants.py new file mode 100644 index 000000000..12c4b4bef --- /dev/null +++ b/comfy/ldm/seedvr/constants.py @@ -0,0 +1,48 @@ +"""SeedVR2 constants.""" + +# Temporal chunk-size law: the sampler's activation wall is linear in +# T_latent * pixel area (17-cell resolution sweep + T bisection, RTX 5090, 3b fp16): +# max_latent_frames = (free_GiB - RESERVED - K*SIGMA) / (GIB_PER_MPX_FRAME * megapixels) +# RESERVED covers model staging plus fixed CUDA/torch overhead; SIGMA is the measured +# run-to-run spread of the wall; K=4 trades ~10% smaller chunks for ~1e-5 OOM odds. +SEEDVR2_CHUNK_GIB_PER_MPX_FRAME = 0.55 +SEEDVR2_CHUNK_RESERVED_GIB = 8.5 +SEEDVR2_CHUNK_SIGMA_GIB = 0.55 +SEEDVR2_CHUNK_SIGMA_K = 4 + +SEEDVR2_7B_VID_DIM = 3072 +SEEDVR2_OOM_BACKOFF_DIVISOR = 2 +SEEDVR2_DTYPE_BYTES_FLOOR = 4 +SEEDVR2_7B_MLP_CHUNK = 8192 +SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS = 4096 # partial-RoPE application token-chunk. +SEEDVR2_LATENT_CHANNELS = 16 + +SEEDVR2_COLOR_MEM_HEADROOM = 0.75 +SEEDVR2_LAB_SCALE_MULTIPLIER = 13 +SEEDVR2_WAVELET_SCALE_MULTIPLIER = 10 # per-frame byte multiplier, wavelet path. +SEEDVR2_ADAIN_SCALE_MULTIPLIER = 6 + +BYTEDANCE_VAE_SCALING_FACTOR = 0.9152 # configs_3b/main.yaml:57. +BYTEDANCE_VAE_SHIFTING_FACTOR = 0.0 +BYTEDANCE_VAE_CONV_MEM_GIB = 0.5 +BYTEDANCE_VAE_NORM_MEM_GIB = 0.5 +BYTEDANCE_LOGVAR_CLAMP_MIN = -30.0 # video_vae_v3/modules/types.py:28. +BYTEDANCE_LOGVAR_CLAMP_MAX = 20.0 # video_vae_v3/modules/types.py:28. +BYTEDANCE_GN_CHUNKS_FP16 = 4 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp16). +BYTEDANCE_GN_CHUNKS_FP32 = 2 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp32). +BYTEDANCE_BLOCK_OUT_CHANNELS = (128, 256, 512, 512) # s8_c16_t4_inflation_sd3.yaml:7-11. +BYTEDANCE_SLICING_SAMPLE_MIN = 4 # s8_c16_t4_inflation_sd3.yaml:22 (slicing_sample_min_size). +BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE = 4 # infer.py:230 (temporal_downsample_factor); the 4n+1 factor. +BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE = 8 # infer.py:231 (spatial_downsample_factor). +BYTEDANCE_720P_REF_AREA = 45 * 80 # dit_v2/window.py:32 (720p reference area for window scaling). +BYTEDANCE_MAX_TEMPORAL_WINDOW = 30 # dit_v2/window.py:35 (max temporal window frames). +BYTEDANCE_ROPE_MAX_FREQ = 256 # dit_v2/rope.py:31 (pixel-RoPE max frequency). +BYTEDANCE_SINUSOIDAL_DIM = 256 # dit_3b/nadit.py:120 (timestep sinusoidal embed dim). + +ROPE_THETA = 10000 # RoPE base; Su et al., "RoFormer", arXiv:2104.09864. + +CIELAB_DELTA = 6.0 / 29.0 # CIE 15 (delta). +CIELAB_KAPPA = (29.0 / 3.0) ** 3 # CIE 15 (kappa). +D65_WHITE_X = 0.95047 # CIE D65 standard illuminant Xn (Yn = 1). +D65_WHITE_Z = 1.08883 # CIE D65 standard illuminant Zn. +WAVELET_DECOMP_LEVELS = 5 # wavelet color-fix decomposition depth (GIMP/Krita; StableSR). diff --git a/comfy/ldm/seedvr/model.py b/comfy/ldm/seedvr/model.py new file mode 100644 index 000000000..a978698d5 --- /dev/null +++ b/comfy/ldm/seedvr/model.py @@ -0,0 +1,1361 @@ +from dataclasses import dataclass +from typing import Optional, Tuple, Union, List, Dict, Any, Callable +import torch.nn.functional as F +from math import ceil, pi +import torch +from itertools import accumulate, chain +from comfy.ldm.modules.diffusionmodules.model import get_timestep_embedding +from comfy.ldm.seedvr.attention import optimized_var_attention +from torch.nn.modules.utils import _triple +from torch import nn +import math +from comfy.ldm.flux.math import apply_rope1 +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_720P_REF_AREA, + BYTEDANCE_MAX_TEMPORAL_WINDOW, + BYTEDANCE_ROPE_MAX_FREQ, + BYTEDANCE_SINUSOIDAL_DIM, + ROPE_THETA, + SEEDVR2_7B_MLP_CHUNK, + SEEDVR2_7B_VID_DIM, + SEEDVR2_LATENT_CHANNELS, + SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS, +) +import comfy.model_management +import comfy.ops + +class Cache: + def __init__(self, disable=False, prefix="", cache=None): + self.cache = cache if cache is not None else {} + self.disable = disable + self.prefix = prefix + + def __call__(self, key: str, fn: Callable): + if self.disable: + return fn() + + key = self.prefix + key + if key not in self.cache: + result = fn() + self.cache[key] = result + return self.cache[key] + + def namespace(self, namespace: str): + return Cache( + disable=self.disable, + prefix=self.prefix + namespace + ".", + cache=self.cache, + ) + +def repeat_concat( + vid: torch.FloatTensor, # (VL ... c) + txt: torch.FloatTensor, # (TL ... c) + vid_len: torch.LongTensor, # (n*b) + txt_len: torch.LongTensor, # (b) + txt_repeat: List, # (n) +) -> torch.FloatTensor: # (L ... c) + vid = torch.split(vid, vid_len.tolist()) + txt = torch.split(txt, txt_len.tolist()) + txt = [[x] * n for x, n in zip(txt, txt_repeat)] + txt = list(chain(*txt)) + return torch.cat(list(chain(*zip(vid, txt)))) + +def repeat_concat_idx( + vid_len: torch.LongTensor, # (n*b) + txt_len: torch.LongTensor, # (b) + txt_repeat: torch.LongTensor, # (n) +) -> Tuple[ + Callable, + Callable, +]: + device = vid_len.device + vid_idx = torch.arange(vid_len.sum(), device=device) + txt_idx = torch.arange(len(vid_idx), len(vid_idx) + txt_len.sum(), device=device) + txt_repeat_list = txt_repeat.tolist() + tgt_idx = repeat_concat(vid_idx, txt_idx, vid_len, txt_len, txt_repeat_list) + src_idx = torch.argsort(tgt_idx) + txt_idx_len = len(tgt_idx) - len(vid_idx) + repeat_txt_len = (txt_len * txt_repeat).tolist() + + def unconcat_coalesce(all): + vid_out, txt_out = all[src_idx].split([len(vid_idx), txt_idx_len]) + txt_out_coalesced = [] + for txt, repeat_time in zip(txt_out.split(repeat_txt_len), txt_repeat_list): + txt = txt.reshape(-1, repeat_time, *txt.shape[1:]).mean(1) + txt_out_coalesced.append(txt) + return vid_out, torch.cat(txt_out_coalesced) + + return ( + lambda vid, txt: torch.cat([vid, txt])[tgt_idx], + lambda all: unconcat_coalesce(all), + ) + +def cumulative_lengths(lengths): + return [0, *accumulate(lengths)] + + +@dataclass +class MMArg: + vid: Any + txt: Any + +def get_args(key: str, args: List[Any]) -> List[Any]: + return [getattr(v, key) if isinstance(v, MMArg) else v for v in args] + + +def get_kwargs(key: str, kwargs: Dict[str, Any]) -> Dict[str, Any]: + return {k: getattr(v, key) if isinstance(v, MMArg) else v for k, v in kwargs.items()} + + +def get_window_op(name: str): + if name == "720pwin_by_size_bysize": + return make_720Pwindows_bysize + if name == "720pswin_by_size_bysize": + return make_shifted_720Pwindows_bysize + raise ValueError(f"Unknown windowing method: {name}") + + +def make_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]): + t, h, w = size + resized_nt, resized_nh, resized_nw = num_windows + scale = math.sqrt(BYTEDANCE_720P_REF_AREA / (h * w)) + resized_h, resized_w = round(h * scale), round(w * scale) + wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) + wt = ceil(min(t, BYTEDANCE_MAX_TEMPORAL_WINDOW) / resized_nt) + nt, nh, nw = ceil(t / wt), ceil(h / wh), ceil(w / ww) + return [ + ( + slice(it * wt, min((it + 1) * wt, t)), + slice(ih * wh, min((ih + 1) * wh, h)), + slice(iw * ww, min((iw + 1) * ww, w)), + ) + for iw in range(nw) + if min((iw + 1) * ww, w) > iw * ww + for ih in range(nh) + if min((ih + 1) * wh, h) > ih * wh + for it in range(nt) + if min((it + 1) * wt, t) > it * wt + ] + +def make_shifted_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]): + t, h, w = size + resized_nt, resized_nh, resized_nw = num_windows + scale = math.sqrt(BYTEDANCE_720P_REF_AREA / (h * w)) + resized_h, resized_w = round(h * scale), round(w * scale) + wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) + wt = ceil(min(t, BYTEDANCE_MAX_TEMPORAL_WINDOW) / resized_nt) + + st, sh, sw = ( + 0.5 if wt < t else 0, + 0.5 if wh < h else 0, + 0.5 if ww < w else 0, + ) + nt, nh, nw = ceil((t - st) / wt), ceil((h - sh) / wh), ceil((w - sw) / ww) + nt, nh, nw = ( + nt + 1 if st > 0 else 1, + nh + 1 if sh > 0 else 1, + nw + 1 if sw > 0 else 1, + ) + return [ + ( + slice(max(int((it - st) * wt), 0), min(int((it - st + 1) * wt), t)), + slice(max(int((ih - sh) * wh), 0), min(int((ih - sh + 1) * wh), h)), + slice(max(int((iw - sw) * ww), 0), min(int((iw - sw + 1) * ww), w)), + ) + for iw in range(nw) + if min(int((iw - sw + 1) * ww), w) > max(int((iw - sw) * ww), 0) + for ih in range(nh) + if min(int((ih - sh + 1) * wh), h) > max(int((ih - sh) * wh), 0) + for it in range(nt) + if min(int((it - st + 1) * wt), t) > max(int((it - st) * wt), 0) + ] + +class RotaryEmbedding(nn.Module): + def __init__( + self, + dim, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + ): + super().__init__() + + self.freqs_for = freqs_for + + if freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + else: + raise ValueError(f"Unknown rotary frequency type: {freqs_for}") + + self.register_buffer("freqs", freqs) + + @property + def device(self): + return self.freqs.device + + def get_axial_freqs( + self, + *dims, + offsets = None + ): + Colon = slice(None) + all_freqs = [] + + if exists(offsets): + if len(offsets) != len(dims): + raise ValueError(f"SeedVR2 rotary offsets length must match dims length, got {len(offsets)} and {len(dims)}.") + + for ind, dim in enumerate(dims): + + offset = 0 + if exists(offsets): + offset = offsets[ind] + + if self.freqs_for == 'pixel': + pos = torch.linspace(-1, 1, steps = dim, device = self.device) + else: + pos = torch.arange(dim, device = self.device) + + pos = pos + offset + + freqs = self.forward(pos) + + all_axis = [None] * len(dims) + all_axis[ind] = Colon + + new_axis_slice = (Ellipsis, *all_axis, Colon) + all_freqs.append(freqs[new_axis_slice]) + + all_freqs = torch.broadcast_tensors(*all_freqs) + return torch.cat(all_freqs, dim = -1) + + def forward( + self, + t, + ): + freqs = self.freqs + + freqs = torch.einsum('..., f -> ... f', t.type(freqs.dtype), freqs) + freqs = freqs.unsqueeze(-1).expand(*freqs.shape, 2).flatten(-2) + + return freqs + +class RotaryEmbeddingBase(nn.Module): + def __init__(self, dim: int, rope_dim: int): + super().__init__() + self.rope = RotaryEmbedding( + dim=dim // rope_dim, + freqs_for="pixel", + max_freq=BYTEDANCE_ROPE_MAX_FREQ, + ) + + def get_axial_freqs(self, *dims): + return self.rope.get_axial_freqs(*dims) + + +class RotaryEmbedding3d(RotaryEmbeddingBase): + def __init__(self, dim: int): + super().__init__(dim, rope_dim=3) + self.mm = False + + +class NaRotaryEmbedding3d(RotaryEmbedding3d): + def forward( + self, + q: torch.FloatTensor, + k: torch.FloatTensor, + shape: torch.LongTensor, + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + freqs = cache("rope_freqs_3d", lambda: self.get_freqs(shape)) + freqs = freqs.to(device=q.device) + q = q.transpose(0, 1) + k = k.transpose(0, 1) + q = _apply_seedvr2_rotary_emb(freqs, q.float()).to(q.dtype) + k = _apply_seedvr2_rotary_emb(freqs, k.float()).to(k.dtype) + q = q.transpose(0, 1) + k = k.transpose(0, 1) + return q, k + + @torch._dynamo.disable + def get_freqs( + self, + shape: torch.LongTensor, + ) -> torch.Tensor: + # Primary provenance: ByteDance-Seed/SeedVR models/dit/rope.py builds + # 7B pixel RoPE with the interleaved-angle convention, not Comfy's + # Flux freqs_cis matrix. + plain_rope = RotaryEmbedding( + dim=self.rope.freqs.numel() * 2, + freqs_for="pixel", + max_freq=BYTEDANCE_ROPE_MAX_FREQ, + ) + plain_rope = plain_rope.to(self.rope.device) + freq_list = [] + for f, h, w in shape.tolist(): + freqs = plain_rope.get_axial_freqs(f, h, w) + freq_list.append(freqs.view(-1, freqs.size(-1))) + return torch.cat(freq_list, dim=0) + + +class MMRotaryEmbeddingBase(RotaryEmbeddingBase): + def __init__(self, dim: int, rope_dim: int): + super().__init__(dim, rope_dim) + self.rope = RotaryEmbedding( + dim=dim // rope_dim, + freqs_for="lang", + theta=ROPE_THETA, + ) + self.mm = True + +def slice_at_dim(t, dim_slice: slice, *, dim): + dim += (t.ndim if dim < 0 else 0) + colons = [slice(None)] * t.ndim + colons[dim] = dim_slice + return t[tuple(colons)] + +def rotate_half(x): + x = x.reshape(*x.shape[:-1], x.shape[-1] // 2, 2) + x1, x2 = x.unbind(dim = -1) + x = torch.stack((-x2, x1), dim = -1) + return x.flatten(-2) +def exists(val): + return val is not None + +def _apply_seedvr2_rotary_emb( + freqs: torch.Tensor, + t: torch.Tensor, + start_index: int = 0, + scale: float = 1.0, + seq_dim: int = -2, + freqs_seq_dim: int | None = None, +) -> torch.Tensor: + dtype = t.dtype + if freqs_seq_dim is None and (freqs.ndim == 2 or t.ndim == 3): + freqs_seq_dim = 0 + + if t.ndim == 3 or freqs_seq_dim is not None: + seq_len = t.shape[seq_dim] + freqs = slice_at_dim(freqs, slice(-seq_len, None), dim=freqs_seq_dim) + + rot_feats = freqs.shape[-1] + end_index = start_index + rot_feats + + t_left = t[..., :start_index] + t_middle = t[..., start_index:end_index] + t_right = t[..., end_index:] + + freqs = freqs.to(device=t_middle.device, dtype=t_middle.dtype) + cos = freqs.cos() * scale + sin = freqs.sin() * scale + t_middle = (t_middle * cos) + (rotate_half(t_middle) * sin) + return torch.cat((t_left, t_middle, t_right), dim=-1).to(dtype) + +def _to_flux_freqs_cis(freqs_interleaved: torch.Tensor) -> torch.Tensor: + angles = freqs_interleaved[..., ::2].float() + cos = torch.cos(angles) + sin = torch.sin(angles) + out = torch.stack([cos, -sin, sin, cos], dim=-1) + return out.reshape(*out.shape[:-1], 2, 2) + + +def _apply_rope1_partial(t: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: + out = t.clone() if t.requires_grad or comfy.model_management.in_training else t + rot_d = 2 * freqs_cis.shape[-3] + seq_len = out.shape[-2] + for start in range(0, seq_len, SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS): + end = min(start + SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS, seq_len) + freqs_chunk = freqs_cis[start:end] + if rot_d == out.shape[-1]: + out[..., start:end, :] = apply_rope1(out[..., start:end, :], freqs_chunk).to(out.dtype) + else: + out[..., start:end, :rot_d] = apply_rope1(out[..., start:end, :rot_d], freqs_chunk).to(out.dtype) + return out + + +class NaMMRotaryEmbedding3d(MMRotaryEmbeddingBase): + def __init__(self, dim: int): + super().__init__(dim, rope_dim=3) + + def forward( + self, + vid_q: torch.FloatTensor, # L h d + vid_k: torch.FloatTensor, # L h d + vid_shape: torch.LongTensor, # B 3 + txt_q: torch.FloatTensor, # L h d + txt_k: torch.FloatTensor, # L h d + txt_shape: torch.LongTensor, # B 1 + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_freqs, txt_freqs = cache( + "mmrope_freqs_3d", + lambda: self.get_freqs(vid_shape, txt_shape), + ) + target_device = vid_q.device + if vid_freqs.device != target_device: + vid_freqs = vid_freqs.to(target_device) + if txt_freqs.device != target_device: + txt_freqs = txt_freqs.to(target_device) + vid_q = vid_q.transpose(0, 1) + vid_k = vid_k.transpose(0, 1) + vid_q = _apply_rope1_partial(vid_q, vid_freqs) + vid_k = _apply_rope1_partial(vid_k, vid_freqs) + vid_q = vid_q.transpose(0, 1) + vid_k = vid_k.transpose(0, 1) + + txt_q = txt_q.transpose(0, 1) + txt_k = txt_k.transpose(0, 1) + txt_q = _apply_rope1_partial(txt_q, txt_freqs) + txt_k = _apply_rope1_partial(txt_k, txt_freqs) + txt_q = txt_q.transpose(0, 1) + txt_k = txt_k.transpose(0, 1) + return vid_q, vid_k, txt_q, txt_k + + @torch._dynamo.disable # Disable compilation: .tolist() is data-dependent and causes graph breaks + def get_freqs( + self, + vid_shape: torch.LongTensor, + txt_shape: torch.LongTensor, + ) -> Tuple[ + torch.Tensor, + torch.Tensor, + ]: + + max_temporal = 0 + max_height = 0 + max_width = 0 + max_txt_len = 0 + + for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()): + max_temporal = max(max_temporal, l + f) + max_height = max(max_height, h) + max_width = max(max_width, w) + max_txt_len = max(max_txt_len, l) + + autocast_device = "cuda" if torch.cuda.is_available() else "cpu" + with torch.amp.autocast(autocast_device, enabled=False): + vid_freqs = self.get_axial_freqs( + max_temporal + 16, + max_height + 4, + max_width + 4, + ).float() + txt_freqs = self.get_axial_freqs(max_txt_len + 16) + + vid_freq_list, txt_freq_list = [], [] + for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()): + vid_freq = vid_freqs[l : l + f, :h, :w].reshape(-1, vid_freqs.size(-1)) + txt_freq = txt_freqs[:l].repeat(1, 3).reshape(-1, vid_freqs.size(-1)) + vid_freq_list.append(vid_freq) + txt_freq_list.append(txt_freq) + vid_freqs_interleaved = torch.cat(vid_freq_list, dim=0) + txt_freqs_interleaved = torch.cat(txt_freq_list, dim=0) + + return _to_flux_freqs_cis(vid_freqs_interleaved), _to_flux_freqs_cis(txt_freqs_interleaved) + +class MMModule(nn.Module): + def __init__( + self, + module: Callable[..., nn.Module], + *args, + shared_weights: bool = False, + vid_only: bool = False, + **kwargs, + ): + super().__init__() + self.shared_weights = shared_weights + self.vid_only = vid_only + if self.shared_weights: + if get_args("vid", args) != get_args("txt", args): + raise ValueError("SeedVR2 shared MMModule requires matching vid/txt args.") + if get_kwargs("vid", kwargs) != get_kwargs("txt", kwargs): + raise ValueError("SeedVR2 shared MMModule requires matching vid/txt kwargs.") + self.all = module(*get_args("vid", args), **get_kwargs("vid", kwargs)) + else: + self.vid = module(*get_args("vid", args), **get_kwargs("vid", kwargs)) + self.txt = ( + module(*get_args("txt", args), **get_kwargs("txt", kwargs)) + if not vid_only + else None + ) + + def forward( + self, + vid: torch.FloatTensor, + txt: torch.FloatTensor, + *args, + **kwargs, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_module = self.vid if not self.shared_weights else self.all + vid = vid_module(vid, *get_args("vid", args), **get_kwargs("vid", kwargs)) + if not self.vid_only: + txt_module = self.txt if not self.shared_weights else self.all + txt = txt.to(device=vid.device, dtype=vid.dtype) + txt = txt_module(txt, *get_args("txt", args), **get_kwargs("txt", kwargs)) + return vid, txt + +def get_na_rope(rope_type: Optional[str], dim: int): + if rope_type is None: + return None + if rope_type == "rope3d": + return NaRotaryEmbedding3d(dim=dim) + if rope_type == "mmrope3d": + return NaMMRotaryEmbedding3d(dim=dim) + raise ValueError(f"Unknown SeedVR2 rope type: {rope_type}") + +class NaMMAttention(nn.Module): + def __init__( + self, + vid_dim: int, + txt_dim: int, + heads: int, + head_dim: int, + qk_bias: bool, + qk_norm, + qk_norm_eps: float, + rope_type: Optional[str], + rope_dim: int, + shared_weights: bool, + device, dtype, operations, + ): + super().__init__() + dim = MMArg(vid_dim, txt_dim) + self.heads = heads + inner_dim = heads * head_dim + qkv_dim = inner_dim * 3 + self.head_dim = head_dim + self.proj_qkv = MMModule( + operations.Linear, dim, qkv_dim, bias=qk_bias, shared_weights=shared_weights, device=device, dtype=dtype + ) + self.proj_out = MMModule(operations.Linear, inner_dim, dim, shared_weights=shared_weights, device=device, dtype=dtype) + self.norm_q = MMModule( + qk_norm, + normalized_shape=head_dim, + eps=qk_norm_eps, + elementwise_affine=True, + shared_weights=shared_weights, + device=device, dtype=dtype + ) + self.norm_k = MMModule( + qk_norm, + normalized_shape=head_dim, + eps=qk_norm_eps, + elementwise_affine=True, + shared_weights=shared_weights, + device=device, dtype=dtype + ) + + + self.rope = get_na_rope(rope_type=rope_type, dim=rope_dim) + +def window( + hid: torch.FloatTensor, # (L c) + hid_shape: torch.LongTensor, # (b n) + window_fn: Callable[[torch.Tensor], List[torch.Tensor]], +): + hid = unflatten(hid, hid_shape) + hid = list(map(window_fn, hid)) + hid_windows_list = [len(x) for x in hid] + hid_windows = torch.as_tensor(hid_windows_list, device=hid_shape.device) + hid = list(chain(*hid)) + hid_len_list = [math.prod(x.shape[:-1]) for x in hid] + hid, hid_shape = flatten(hid) + return hid, hid_shape, hid_windows, hid_len_list, hid_windows_list + +def window_idx( + hid_shape: torch.LongTensor, # (b n) + window_fn: Callable[[torch.Tensor], List[torch.Tensor]], +): + hid_idx = torch.arange(hid_shape.prod(-1).sum(), device=hid_shape.device).unsqueeze(-1) + tgt_idx, tgt_shape, tgt_windows, tgt_len_list, tgt_windows_list = window(hid_idx, hid_shape, window_fn) + tgt_idx = tgt_idx.squeeze(-1) + src_idx = torch.argsort(tgt_idx) + return ( + lambda hid: torch.index_select(hid, 0, tgt_idx), + lambda hid: torch.index_select(hid, 0, src_idx), + tgt_shape, + tgt_windows, + tgt_len_list, + tgt_windows_list, + ) + +class NaSwinAttention(NaMMAttention): + def __init__( + self, + *args, + window: Union[int, Tuple[int, int, int]], + window_method: str, + version: bool = False, + **kwargs, + ): + super().__init__(*args, **kwargs) + self.version_7b = version + self.window = _triple(window) + self.window_method = window_method + if not all(isinstance(v, int) and v >= 0 for v in self.window): + raise ValueError(f"SeedVR2 window must contain non-negative integers, got {self.window}.") + + self.window_op = get_window_op(window_method) + + def forward( + self, + vid: torch.FloatTensor, # l c + txt: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, # b 3 + txt_shape: torch.LongTensor, # b 1 + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + + vid_qkv, txt_qkv = self.proj_qkv(vid, txt) + + cache_win = cache.namespace(f"{self.window_method}_{self.window}_sd3") + + def make_window(x: torch.Tensor): + t, h, w, _ = x.shape + window_slices = self.window_op((t, h, w), self.window) + return [x[st, sh, sw] for (st, sh, sw) in window_slices] + + window_partition, window_reverse, window_shape, window_count, vid_len_win_list, window_count_list = cache_win( + "win_transform", + lambda: window_idx(vid_shape, make_window), + ) + vid_qkv_win = window_partition(vid_qkv) + + vid_qkv_win = vid_qkv_win.reshape(vid_qkv_win.shape[0], 3, self.heads, self.head_dim) + txt_qkv = txt_qkv.reshape(txt_qkv.shape[0], 3, self.heads, self.head_dim) + + vid_q, vid_k, vid_v = vid_qkv_win.unbind(1) + txt_q, txt_k, txt_v = txt_qkv.unbind(1) + + vid_q, txt_q = self.norm_q(vid_q, txt_q) + vid_k, txt_k = self.norm_k(vid_k, txt_k) + + txt_len = cache("txt_len", lambda: txt_shape.prod(-1)) + + vid_len_win = cache_win("vid_len", lambda: window_shape.prod(-1)) + txt_len = txt_len.to(window_count.device) + + if self.rope: + if self.version_7b: + vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + elif self.rope.mm: + _, num_h, _ = txt_q.shape + txt_q_repeat = txt_q.flatten(1, 2) + txt_q_repeat = unflatten(txt_q_repeat, txt_shape) + txt_q_repeat = [[x] * n for x, n in zip(txt_q_repeat, window_count_list)] + txt_q_repeat = list(chain(*txt_q_repeat)) + txt_q_repeat, txt_shape_repeat = flatten(txt_q_repeat) + txt_q_repeat = txt_q_repeat.reshape(txt_q_repeat.shape[0], num_h, self.head_dim) + + txt_k_repeat = txt_k.flatten(1, 2) + txt_k_repeat = unflatten(txt_k_repeat, txt_shape) + txt_k_repeat = [[x] * n for x, n in zip(txt_k_repeat, window_count_list)] + txt_k_repeat = list(chain(*txt_k_repeat)) + txt_k_repeat, _ = flatten(txt_k_repeat) + txt_k_repeat = txt_k_repeat.reshape(txt_k_repeat.shape[0], num_h, self.head_dim) + + vid_q, vid_k, txt_q, txt_k = self.rope( + vid_q, vid_k, window_shape, txt_q_repeat, txt_k_repeat, txt_shape_repeat, cache_win + ) + else: + vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + + txt_len_win_list = cache_win( + "txt_len_list", + lambda: [txt_len for txt_len, window_count in zip(txt_len.tolist(), window_count_list) for _ in range(window_count)], + ) + all_len_win = cache_win("all_len", lambda: [vid_len + txt_len for vid_len, txt_len in zip(vid_len_win_list, txt_len_win_list)]) + concat_win, unconcat_win = cache_win( + "mm_pnp", lambda: repeat_concat_idx(vid_len_win, txt_len, window_count) + ) + out = optimized_var_attention( + q=concat_win(vid_q, txt_q), + k=concat_win(vid_k, txt_k), + v=concat_win(vid_v, txt_v), + heads=self.heads, skip_reshape=True, skip_output_reshape=True, + cu_seqlens_q=cache_win("vid_seqlens_q", lambda: cumulative_lengths(all_len_win)), + cu_seqlens_k=cache_win("vid_seqlens_k", lambda: cumulative_lengths(all_len_win)), + ) + vid_out, txt_out = unconcat_win(out) + + vid_out = vid_out.flatten(1, 2) + txt_out = txt_out.flatten(1, 2) + vid_out = window_reverse(vid_out) + + vid_out, txt_out = self.proj_out(vid_out, txt_out) + + return vid_out, txt_out + +class MLP(nn.Module): + def __init__( + self, + dim: int, + expand_ratio: int, + device, dtype, operations + ): + super().__init__() + self.proj_in = operations.Linear(dim, dim * expand_ratio, device=device, dtype=dtype) + self.act = nn.GELU("tanh") + self.proj_out = operations.Linear(dim * expand_ratio, dim, device=device, dtype=dtype) + + def forward(self, x: torch.FloatTensor) -> torch.FloatTensor: + x = self.proj_in(x) + x = self.act(x) + x = self.proj_out(x) + return x + + +class SwiGLUMLP(nn.Module): + def __init__( + self, + dim: int, + expand_ratio: int, + multiple_of: int = 256, + device=None, dtype=None, operations=None + ): + super().__init__() + hidden_dim = int(2 * dim * expand_ratio / 3) + hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of) + self.proj_in_gate = operations.Linear(dim, hidden_dim, bias=False, device=device, dtype=dtype) + self.proj_out = operations.Linear(hidden_dim, dim, bias=False, device=device, dtype=dtype) + self.proj_in = operations.Linear(dim, hidden_dim, bias=False, device=device, dtype=dtype) + + def forward(self, x: torch.FloatTensor) -> torch.FloatTensor: + return self.proj_out(F.silu(self.proj_in_gate(x)) * self.proj_in(x)) + +def get_mlp(mlp_type: Optional[str] = "normal"): + if mlp_type == "normal": + return MLP + if mlp_type == "swiglu": + return SwiGLUMLP + raise ValueError(f"Unknown SeedVR2 MLP type: {mlp_type}") + +class NaMMSRTransformerBlock(nn.Module): + def __init__( + self, + *, + vid_dim: int, + txt_dim: int, + emb_dim: int, + heads: int, + head_dim: int, + expand_ratio: int, + norm, + norm_eps: float, + ada, + qk_bias: bool, + qk_norm, + mlp_type: str, + shared_weights: bool, + rope_type: str, + rope_dim: int, + is_last_layer: bool, + window: Union[int, Tuple[int, int, int]], + window_method: str, + version: bool, + device, dtype, operations, + ): + super().__init__() + dim = MMArg(vid_dim, txt_dim) + self.attn_norm = MMModule(norm, normalized_shape=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights, device=device, dtype=dtype) + + self.attn = NaSwinAttention( + vid_dim=vid_dim, + txt_dim=txt_dim, + heads=heads, + head_dim=head_dim, + qk_bias=qk_bias, + qk_norm=qk_norm, + qk_norm_eps=norm_eps, + rope_type=rope_type, + rope_dim=rope_dim, + shared_weights=shared_weights, + window=window, + window_method=window_method, + version=version, + device=device, dtype=dtype, operations=operations + ) + + self.mlp_norm = MMModule(norm, normalized_shape=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights, vid_only=is_last_layer, device=device, dtype=dtype) + self.mlp = MMModule( + get_mlp(mlp_type), + dim=dim, + expand_ratio=expand_ratio, + shared_weights=shared_weights, + vid_only=is_last_layer, + device=device, dtype=dtype, operations=operations + ) + self.ada = MMModule(ada, dim=dim, emb_dim=emb_dim, layers=["attn", "mlp"], shared_weights=shared_weights, vid_only=is_last_layer, device=device, dtype=dtype) + self.is_last_layer = is_last_layer + self.version = version + + def _seedvr2_7b_mlp( + self, + vid: torch.FloatTensor, + txt: torch.FloatTensor, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + ]: + vid_module = self.mlp.vid if not self.mlp.shared_weights else self.mlp.all + if comfy.model_management.in_training or vid.requires_grad: + vid = torch.cat([vid_module(chunk) for chunk in vid.split(SEEDVR2_7B_MLP_CHUNK, dim=0)], dim=0) + else: + vid_out = None + offset = 0 + for chunk in vid.split(SEEDVR2_7B_MLP_CHUNK, dim=0): + chunk_out = vid_module(chunk) + if vid_out is None: + vid_out = chunk_out.new_empty((vid.shape[0], *chunk_out.shape[1:])) + vid_out[offset:offset + chunk_out.shape[0]] = chunk_out + offset += chunk_out.shape[0] + vid = vid_out + if not self.mlp.vid_only: + txt_module = self.mlp.txt if not self.mlp.shared_weights else self.mlp.all + txt = txt.to(device=vid.device, dtype=vid.dtype) + txt = txt_module(txt) + return vid, txt + + def forward( + self, + vid: torch.FloatTensor, # l c + txt: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, # b 3 + txt_shape: torch.LongTensor, # b 1 + emb: torch.FloatTensor, + cache: Cache, + ) -> Tuple[ + torch.FloatTensor, + torch.FloatTensor, + torch.LongTensor, + torch.LongTensor, + ]: + hid_len = MMArg( + cache("vid_len", lambda: vid_shape.prod(-1)), + cache("txt_len", lambda: txt_shape.prod(-1)), + ) + ada_kwargs = { + "emb": emb, + "hid_len": hid_len, + "cache": cache, + "branch_tag": MMArg("vid", "txt"), + } + + vid_attn, txt_attn = self.attn_norm(vid, txt) + vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="in", **ada_kwargs) + vid_attn, txt_attn = self.attn(vid_attn, txt_attn, vid_shape, txt_shape, cache) + vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="out", **ada_kwargs) + vid_attn, txt_attn = (vid_attn + vid), (txt_attn + txt) + + vid_mlp, txt_mlp = self.mlp_norm(vid_attn, txt_attn) + vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="in", **ada_kwargs) + if self.version: + vid_mlp, txt_mlp = self._seedvr2_7b_mlp(vid_mlp, txt_mlp) + else: + vid_mlp, txt_mlp = self.mlp(vid_mlp, txt_mlp) + vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="out", **ada_kwargs) + vid_mlp, txt_mlp = (vid_mlp + vid_attn), (txt_mlp + txt_attn) + + return vid_mlp, txt_mlp, vid_shape, txt_shape + +class PatchOut(nn.Module): + def __init__( + self, + out_channels: int, + patch_size: Union[int, Tuple[int, int, int]], + dim: int, + device, dtype, operations + ): + super().__init__() + t, h, w = _triple(patch_size) + self.patch_size = t, h, w + self.proj = operations.Linear(dim, out_channels * t * h * w, device=device, dtype=dtype) + + def forward( + self, + vid: torch.Tensor, + ) -> torch.Tensor: + t, h, w = self.patch_size + vid = self.proj(vid) + b, T, H, W, channels = vid.shape + c = channels // (t * h * w) + vid = vid.view(b, T, H, W, t, h, w, c).permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(b, c, T * t, H * h, W * w) + if t > 1: + vid = vid[:, :, (t - 1) :] + return vid + +class NaPatchOut(PatchOut): + def forward( + self, + vid: torch.FloatTensor, # l c + vid_shape: torch.LongTensor, + cache: Optional[Cache] = None, + vid_shape_before_patchify = None + ) -> Tuple[ + torch.FloatTensor, + torch.LongTensor, + ]: + if cache is None: + cache = Cache(disable=True) + + t, h, w = self.patch_size + vid = self.proj(vid) + + if not (t == h == w == 1): + vid = unflatten(vid, vid_shape) + for i in range(len(vid)): + T, H, W, channels = vid[i].shape + c = channels // (t * h * w) + vid[i] = vid[i].view(T, H, W, t, h, w, c).permute(0, 3, 1, 4, 2, 5, 6).reshape(T * t, H * h, W * w, c) + if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: + vid[i] = vid[i][(t - vid_shape_before_patchify[i, 0] % t) :] + vid, vid_shape = flatten(vid) + + return vid, vid_shape + +class PatchIn(nn.Module): + def __init__( + self, + in_channels: int, + patch_size: Union[int, Tuple[int, int, int]], + dim: int, + device, dtype, operations + ): + super().__init__() + t, h, w = _triple(patch_size) + self.patch_size = t, h, w + self.proj = operations.Linear(in_channels * t * h * w, dim, device=device, dtype=dtype) + + def forward( + self, + vid: torch.Tensor, + ) -> torch.Tensor: + t, h, w = self.patch_size + if t > 1: + if vid.size(2) % t != 1: + raise ValueError( + f"SeedVR2 patch input temporal size must satisfy T % {t} == 1, got {vid.size(2)}." + ) + vid = torch.cat([vid[:, :, :1]] * (t - 1) + [vid], dim=2) + b, c, Tt, Hh, Ww = vid.shape + vid = vid.view(b, c, Tt // t, t, Hh // h, h, Ww // w, w).permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(b, Tt // t, Hh // h, Ww // w, t * h * w * c) + vid = self.proj(vid) + return vid + +class NaPatchIn(PatchIn): + def forward( + self, + vid: torch.Tensor, # l c + vid_shape: torch.LongTensor, + cache: Optional[Cache] = None, + ) -> torch.Tensor: + if cache is None: + cache = Cache(disable=True) + cache = cache.namespace("patch") + vid_shape_before_patchify = cache("vid_shape_before_patchify", lambda: vid_shape) + t, h, w = self.patch_size + if not (t == h == w == 1): + vid = unflatten(vid, vid_shape) + for i in range(len(vid)): + if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: + vid[i] = torch.cat([vid[i][:1]] * (t - vid[i].size(0) % t) + [vid[i]], dim=0) + Tt, Hh, Ww, c = vid[i].shape + vid[i] = vid[i].view(Tt // t, t, Hh // h, h, Ww // w, w, c).permute(0, 2, 4, 1, 3, 5, 6).reshape(Tt // t, Hh // h, Ww // w, t * h * w * c) + vid, vid_shape = flatten(vid) + + vid = self.proj(vid) + return vid, vid_shape + +def expand_dims(x: torch.Tensor, dim: int, ndim: int): + shape = x.shape + shape = shape[:dim] + (1,) * (ndim - len(shape)) + shape[dim:] + return x.reshape(shape) + + +class AdaSingle(nn.Module): + def __init__( + self, + dim: int, + emb_dim: int, + layers: List[str], + modes: Tuple[str, ...] = ("in", "out"), + device = None, dtype = None, + ): + if emb_dim != 6 * dim: + raise ValueError(f"SeedVR2 AdaSingle requires emb_dim == 6 * dim, got emb_dim={emb_dim}, dim={dim}.") + super().__init__() + self.dim = dim + self.emb_dim = emb_dim + self.layers = layers + + param_kwargs = {"device": device, "dtype": dtype} + + for l in layers: + if "in" in modes: + self.register_parameter(f"{l}_shift", nn.Parameter(torch.empty(dim, **param_kwargs))) + self.register_parameter(f"{l}_scale", nn.Parameter(torch.empty(dim, **param_kwargs))) + if "out" in modes: + self.register_parameter(f"{l}_gate", nn.Parameter(torch.empty(dim, **param_kwargs))) + + def forward( + self, + hid: torch.FloatTensor, # b ... c + emb: torch.FloatTensor, # b d + layer: str, + mode: str, + cache: Optional[Cache] = None, + branch_tag: str = "", + hid_len: Optional[torch.LongTensor] = None, # b + ) -> torch.FloatTensor: + if cache is None: + cache = Cache(disable=True) + idx = self.layers.index(layer) + emb = emb.reshape(emb.shape[0], -1, len(self.layers), 3)[:, :, idx, :] + emb = expand_dims(emb, 1, hid.ndim + 1) + + if hid_len is not None: + emb = cache( + f"emb_repeat_{idx}_{branch_tag}", + lambda: torch.repeat_interleave(emb, hid_len, dim=0), + ) + + shiftA, scaleA, gateA = emb.unbind(-1) + shiftB, scaleB, gateB = ( + getattr(self, f"{layer}_shift", None), + getattr(self, f"{layer}_scale", None), + getattr(self, f"{layer}_gate", None), + ) + + if mode == "in": + shiftB = comfy.ops.cast_to_input(shiftB, hid) + scaleB = comfy.ops.cast_to_input(scaleB, hid) + return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) + if mode == "out": + if gateB is not None: + gateB = comfy.ops.cast_to_input(gateB, hid) + return hid.mul_(gateA + gateB) + else: + return hid.mul_(gateA) + + raise ValueError(f"Unknown AdaSingle mode: {mode}") + + +class TimeEmbedding(nn.Module): + def __init__( + self, + sinusoidal_dim: int, + hidden_dim: int, + output_dim: int, + device, dtype, operations + ): + super().__init__() + self.sinusoidal_dim = sinusoidal_dim + self.proj_in = operations.Linear(sinusoidal_dim, hidden_dim, device=device, dtype=dtype) + self.proj_hid = operations.Linear(hidden_dim, hidden_dim, device=device, dtype=dtype) + self.proj_out = operations.Linear(hidden_dim, output_dim, device=device, dtype=dtype) + self.act = nn.SiLU() + + def forward( + self, + timestep: Union[int, float, torch.IntTensor, torch.FloatTensor], + device: torch.device, + dtype: torch.dtype, + ) -> torch.FloatTensor: + if not torch.is_tensor(timestep): + timestep = torch.tensor([timestep], device=device, dtype=dtype) + if timestep.ndim == 0: + timestep = timestep[None] + + emb = get_timestep_embedding( + timesteps=timestep, + embedding_dim=self.sinusoidal_dim, + flip_sin_to_cos=False, + downscale_freq_shift=0, + ).to(dtype) + emb = self.proj_in(emb) + emb = self.act(emb) + emb = self.proj_hid(emb) + emb = self.act(emb) + emb = self.proj_out(emb) + return emb + +def flatten( + hid: List[torch.FloatTensor], # List of (*** c) +) -> Tuple[ + torch.FloatTensor, # (L c) + torch.LongTensor, # (b n) +]: + if len(hid) == 0: + raise ValueError("SeedVR2 flatten requires at least one tensor.") + shape = torch.as_tensor([x.shape[:-1] for x in hid], device=hid[0].device) + hid = torch.cat([x.flatten(0, -2) for x in hid]) + return hid, shape + + +def unflatten( + hid: torch.FloatTensor, # (L c) or (L ... c) + hid_shape: torch.LongTensor, # (b n) +) -> List[torch.Tensor]: # List of (*** c) or (*** ... c) + hid_len = hid_shape.prod(-1) + hid = hid.split(hid_len.tolist()) + hid = [x.unflatten(0, s.tolist()) for x, s in zip(hid, hid_shape)] + return hid + +class NaDiT(nn.Module): + + def __init__( + self, + norm_eps, + num_layers, + mlp_type, + vid_in_channels = 33, + vid_out_channels = SEEDVR2_LATENT_CHANNELS, + vid_dim = 2560, + txt_in_dim = 5120, + heads = 20, + head_dim = 128, + mm_layers = 10, + expand_ratio = 4, + qk_bias = False, + patch_size = (1, 2, 2), + rope_dim = 128, + rope_type = "mmrope3d", + vid_out_norm: Optional[str] = None, + image_model = None, + device = None, + dtype = None, + operations = None, + ): + if image_model not in (None, "seedvr2"): + raise ValueError(f"SeedVR2 NaDiT expected image_model='seedvr2', got {image_model!r}.") + self._7b_version = vid_dim == SEEDVR2_7B_VID_DIM + if self._7b_version: + rope_type = "rope3d" + self.dtype = dtype + factory_kwargs = {"device": device, "dtype": dtype} + window_method = num_layers // 2 * ["720pwin_by_size_bysize","720pswin_by_size_bysize"] + txt_dim = vid_dim + emb_dim = vid_dim * 6 + window = num_layers * [(4,3,3)] + ada = AdaSingle + norm = operations.RMSNorm + qk_norm = operations.RMSNorm + super().__init__() + self.register_buffer("positive_conditioning", torch.empty((58, 5120), device=device, dtype=dtype)) + self.register_buffer("negative_conditioning", torch.empty((64, 5120), device=device, dtype=dtype)) + self.vid_in = NaPatchIn( + in_channels=vid_in_channels, + patch_size=patch_size, + dim=vid_dim, + device=device, dtype=dtype, operations=operations + ) + self.txt_in = ( + operations.Linear(txt_in_dim, txt_dim, **factory_kwargs) + if txt_in_dim and txt_in_dim != txt_dim + else nn.Identity() + ) + self.emb_in = TimeEmbedding( + sinusoidal_dim=BYTEDANCE_SINUSOIDAL_DIM, + hidden_dim=max(vid_dim, txt_dim), + output_dim=emb_dim, + device=device, dtype=dtype, operations=operations + ) + + if window is None or isinstance(window[0], int): + window = [window] * num_layers + + rope_dim = rope_dim if rope_dim is not None else head_dim // 2 + self.blocks = nn.ModuleList( + [ + NaMMSRTransformerBlock( + vid_dim=vid_dim, + txt_dim=txt_dim, + emb_dim=emb_dim, + heads=heads, + head_dim=head_dim, + expand_ratio=expand_ratio, + norm=norm, + norm_eps=norm_eps, + ada=ada, + qk_bias=qk_bias, + qk_norm=qk_norm, + mlp_type=mlp_type, + rope_dim = rope_dim, + window=window[i], + window_method=window_method[i], + version = self._7b_version, + is_last_layer=(i == num_layers - 1) and not self._7b_version, + rope_type = rope_type, + shared_weights=not ( + (i < mm_layers) if isinstance(mm_layers, int) else mm_layers[i] + ), + operations = operations, + **factory_kwargs + ) + for i in range(num_layers) + ] + ) + self.vid_out = NaPatchOut( + out_channels=vid_out_channels, + patch_size=patch_size, + dim=vid_dim, + device=device, dtype=dtype, operations=operations + ) + + self.vid_out_norm = None + if vid_out_norm is not None: + self.vid_out_norm = operations.RMSNorm( + normalized_shape=vid_dim, + eps=norm_eps, + elementwise_affine=True, + device=device, dtype=dtype + ) + self.vid_out_ada = ada( + dim=vid_dim, + emb_dim=emb_dim, + layers=["out"], + modes=["in"], + device=device, dtype=dtype + ) + + def _resolve_text_conditioning(self, context, cond_or_uncond=None): + if context is None or context.numel() == 0: + context = self.positive_conditioning + return flatten([context]) + if NaDiT._seedvr2_is_single_conditioning_branch(cond_or_uncond): + if context.shape[0] == 1: + context = context.squeeze(0) + return flatten([context]) + return flatten(context.unbind(0)) + if context.shape[0] % 2 != 0: + raise ValueError(f"SeedVR2 expected an even text-conditioning batch, got shape {tuple(context.shape)}") + neg_cond, pos_cond = context.chunk(2, dim=0) + if pos_cond.shape[0] == 1: + pos_cond, neg_cond = pos_cond.squeeze(0), neg_cond.squeeze(0) + return flatten([pos_cond, neg_cond]) + return flatten((*pos_cond.unbind(0), *neg_cond.unbind(0))) + + @staticmethod + def _seedvr2_is_single_conditioning_branch(cond_or_uncond): + if cond_or_uncond is None or len(cond_or_uncond) == 0: + return False + first = cond_or_uncond[0] + return all(entry == first for entry in cond_or_uncond) + + @staticmethod + def _check_seedvr2_video_latent(x, channels, name): + if x.ndim != 5: + raise ValueError(f"SeedVR2 expected {name} to be 5-D native latent, got shape {tuple(x.shape)}.") + if x.shape[1] != channels: + raise ValueError(f"SeedVR2 expected {name} channels to be {channels}, got shape {tuple(x.shape)}.") + return x + + def _swap_pos_neg_halves(self, out, cond_or_uncond=None): + if NaDiT._seedvr2_is_single_conditioning_branch(cond_or_uncond): + return out + pos, neg = out.chunk(2, dim=0) + return torch.cat([neg, pos], dim=0) + + def forward( + self, + x, + timestep, + context, # l c + disable_cache: bool = False, + **kwargs + ): + transformer_options = kwargs.get("transformer_options", {}) + patches_replace = transformer_options.get("patches_replace", {}) + blocks_replace = patches_replace.get("dit", {}) + conditions = kwargs.get("condition") + if conditions is None: + raise ValueError("SeedVR2 requires conditioning latents from the SeedVR2Conditioning node.") + x = self._check_seedvr2_video_latent(x, SEEDVR2_LATENT_CHANNELS, "latent") + conditions = self._check_seedvr2_video_latent(conditions, SEEDVR2_LATENT_CHANNELS + 1, "conditioning") + b, _, t, h, w = x.shape + if conditions.shape[0] != b or conditions.shape[2:] != (t, h, w): + raise ValueError( + f"SeedVR2 conditioning shape must match latent batch/temporal/spatial dimensions; got latent {tuple(x.shape)} and conditioning {tuple(conditions.shape)}." + ) + x = x.movedim(1, -1) + conditions = conditions.movedim(1, -1) + cache = Cache(disable=disable_cache) + + txt, txt_shape = self._resolve_text_conditioning(context, transformer_options.get("cond_or_uncond")) + + vid, vid_shape = flatten(x) + cond_latent, _ = flatten(conditions) + + vid = torch.cat([vid, cond_latent], dim=-1) + + txt = self.txt_in(txt) + + vid_shape_before_patchify = vid_shape + vid, vid_shape = self.vid_in(vid, vid_shape, cache=cache) + + emb = self.emb_in(timestep, device=vid.device, dtype=vid.dtype) + + for i, block in enumerate(self.blocks): + if ("block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["vid"], out["txt"], out["vid_shape"], out["txt_shape"] = block( + vid=args["vid"], + txt=args["txt"], + vid_shape=args["vid_shape"], + txt_shape=args["txt_shape"], + emb=args["emb"], + cache=args["cache"], + ) + return out + out = blocks_replace[("block", i)]({ + "vid":vid, + "txt":txt, + "vid_shape":vid_shape, + "txt_shape":txt_shape, + "emb":emb, + "cache":cache, + }, {"original_block": block_wrap}) + vid, txt, vid_shape, txt_shape = out["vid"], out["txt"], out["vid_shape"], out["txt_shape"] + else: + vid, txt, vid_shape, txt_shape = block( + vid=vid, + txt=txt, + vid_shape=vid_shape, + txt_shape=txt_shape, + emb=emb, + cache=cache, + ) + + if self.vid_out_norm: + vid = self.vid_out_norm(vid) + vid = self.vid_out_ada( + vid, + emb=emb, + layer="out", + mode="in", + hid_len=cache("vid_len", lambda: vid_shape.prod(-1)), + cache=cache, + branch_tag="vid", + ) + + vid, vid_shape = self.vid_out(vid, vid_shape, cache, vid_shape_before_patchify = vid_shape_before_patchify) + vid = unflatten(vid, vid_shape) + out = torch.stack(vid) + out = out.movedim(-1, 1) + return self._swap_pos_neg_halves(out, transformer_options.get("cond_or_uncond")) diff --git a/comfy/ldm/seedvr/vae.py b/comfy/ldm/seedvr/vae.py new file mode 100644 index 000000000..7a8070b65 --- /dev/null +++ b/comfy/ldm/seedvr/vae.py @@ -0,0 +1,1610 @@ +from typing import Literal, Optional, Tuple +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from contextlib import contextmanager +from comfy.utils import ProgressBar + +from comfy.ldm.seedvr.constants import ( + BYTEDANCE_BLOCK_OUT_CHANNELS, + BYTEDANCE_GN_CHUNKS_FP16, + BYTEDANCE_GN_CHUNKS_FP32, + BYTEDANCE_LOGVAR_CLAMP_MAX, + BYTEDANCE_LOGVAR_CLAMP_MIN, + BYTEDANCE_SLICING_SAMPLE_MIN, + BYTEDANCE_VAE_CONV_MEM_GIB, + BYTEDANCE_VAE_NORM_MEM_GIB, + BYTEDANCE_VAE_SCALING_FACTOR, + BYTEDANCE_VAE_SHIFTING_FACTOR, + BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE, + BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE, + SEEDVR2_LATENT_CHANNELS, +) +from comfy.ldm.modules.attention import optimized_attention +from comfy.ldm.modules.diffusionmodules.model import vae_attention + +import math +from enum import Enum + +import logging +import comfy.model_management +import comfy.ops +ops = comfy.ops.manual_cast + + +def _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap, temporal_scale=1): + if temporal_size is None: + return None + + temporal_size = int(temporal_size) + if temporal_size <= 0: + return None + + temporal_overlap = max(0, int(temporal_overlap or 0)) + temporal_overlap = min(temporal_overlap, temporal_size - 1) + temporal_step = temporal_size - temporal_overlap + temporal_scale = max(1, int(temporal_scale)) + return max(1, math.ceil(temporal_step / temporal_scale)) + + +def _seedvr2_clamped_spatial_overlap(overlap, tile_size): + overlap = max(0, int(overlap)) + tile_size = max(1, int(tile_size)) + return min(overlap, tile_size - 1) + + +def tiled_vae( + x, + vae_model, + tile_size=(512, 512), + tile_overlap=(64, 64), + temporal_size=16, + temporal_overlap=0, + encode=True, +): + if x.ndim != 5: + x = x.unsqueeze(2) + + _, _, d, h, w = x.shape + + sf_s = getattr(vae_model, "spatial_downsample_factor", BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE) + sf_t = getattr(vae_model, "temporal_downsample_factor", BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE) + if encode: + slicing_attr = "slicing_sample_min_size" + slicing_min_size = _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap) + else: + slicing_attr = "slicing_latent_min_size" + slicing_min_size = _seedvr2_temporal_slicing_min_size(temporal_size, temporal_overlap, sf_t) + if encode: + ti_h, ti_w = tile_size + ov_h = _seedvr2_clamped_spatial_overlap(tile_overlap[0], ti_h) + ov_w = _seedvr2_clamped_spatial_overlap(tile_overlap[1], ti_w) + blend_ov_h = max(0, ov_h // sf_s) + blend_ov_w = max(0, ov_w // sf_s) + target_d = (d + sf_t - 1) // sf_t + target_h = (h + sf_s - 1) // sf_s + target_w = (w + sf_s - 1) // sf_s + else: + ti_h = max(1, tile_size[0] // sf_s) + ti_w = max(1, tile_size[1] // sf_s) + ov_h = _seedvr2_clamped_spatial_overlap(tile_overlap[0] // sf_s, ti_h) + ov_w = _seedvr2_clamped_spatial_overlap(tile_overlap[1] // sf_s, ti_w) + blend_ov_h = ov_h * sf_s + blend_ov_w = ov_w * sf_s + + target_d = max(1, d * sf_t - (sf_t - 1)) + target_h = h * sf_s + target_w = w * sf_s + + stride_h = max(1, ti_h - ov_h) + stride_w = max(1, ti_w - ov_w) + + storage_device = vae_model.device + result = None + count = None + def run_temporal_chunks(spatial_tile, model=vae_model): + t_chunk = spatial_tile.contiguous() + old_device = getattr(model, "device", None) + model.device = t_chunk.device + old_slicing_min_size = getattr(model, slicing_attr, None) + if old_slicing_min_size is not None and slicing_min_size is not None: + if slicing_min_size <= 0: + setattr(model, slicing_attr, t_chunk.shape[2]) + else: + setattr(model, slicing_attr, slicing_min_size) + try: + if encode: + out = model.encode(t_chunk) + else: + out = model.decode_(t_chunk) + finally: + if old_slicing_min_size is not None and slicing_min_size is not None: + setattr(model, slicing_attr, old_slicing_min_size) + if old_device is not None: + model.device = old_device + if out.ndim == 4: + out = out.unsqueeze(2) + return out.to(storage_device) + + ramp_cache = {} + def get_ramp(steps): + if steps not in ramp_cache: + t = torch.linspace(0, 1, steps=steps, device=storage_device, dtype=torch.float32) + ramp_cache[steps] = 0.5 - 0.5 * torch.cos(t * torch.pi) + return ramp_cache[steps] + + tile_ranges = [] + for y_idx in range(0, h, stride_h): + y_end = min(y_idx + ti_h, h) + if y_idx > 0 and (y_end - y_idx) <= ov_h: + continue + for x_idx in range(0, w, stride_w): + x_end = min(x_idx + ti_w, w) + if x_idx > 0 and (x_end - x_idx) <= ov_w: + continue + tile_ranges.append((y_idx, y_end, x_idx, x_end)) + + total_tiles = len(tile_ranges) + bar = ProgressBar(total_tiles) + single_spatial_tile = h <= ti_h and w <= ti_w + + def run_tile(tile_index, tile_range): + y_idx, y_end, x_idx, x_end = tile_range + tile_x = x[:, :, :, y_idx:y_end, x_idx:x_end] + tile_out = run_temporal_chunks(tile_x) + return tile_index, y_idx, y_end, x_idx, x_end, tile_out + + ordered_tile_outputs = ( + run_tile(tile_index, tile_range) + for tile_index, tile_range in enumerate(tile_ranges) + ) + + for _, y_idx, y_end, x_idx, x_end, tile_out in ordered_tile_outputs: + + if single_spatial_tile: + result = tile_out[:, :, :target_d, :target_h, :target_w] + if result.device != x.device or result.dtype != x.dtype: + result = result.to(device=x.device, dtype=x.dtype) + if x.shape[2] == 1 and sf_t == 1: + result = result.squeeze(2) + bar.update(1) + return result + + if result is None: + b_out, c_out = tile_out.shape[0], tile_out.shape[1] + result = torch.zeros((b_out, c_out, target_d, target_h, target_w), device=storage_device, dtype=torch.float32) + count = torch.zeros((1, 1, 1, target_h, target_w), device=storage_device, dtype=torch.float32) + + if encode: + ys, ye = y_idx // sf_s, (y_idx // sf_s) + tile_out.shape[3] + xs, xe = x_idx // sf_s, (x_idx // sf_s) + tile_out.shape[4] + cur_ov_h = max(0, min(blend_ov_h, tile_out.shape[3] // 2)) + cur_ov_w = max(0, min(blend_ov_w, tile_out.shape[4] // 2)) + else: + ys, ye = y_idx * sf_s, (y_idx * sf_s) + tile_out.shape[3] + xs, xe = x_idx * sf_s, (x_idx * sf_s) + tile_out.shape[4] + cur_ov_h = max(0, min(blend_ov_h, tile_out.shape[3] // 2)) + cur_ov_w = max(0, min(blend_ov_w, tile_out.shape[4] // 2)) + + w_h = torch.ones((tile_out.shape[3],), device=storage_device) + w_w = torch.ones((tile_out.shape[4],), device=storage_device) + + if cur_ov_h > 0: + r = get_ramp(cur_ov_h) + if y_idx > 0: + w_h[:cur_ov_h] = r + if y_end < h: + w_h[-cur_ov_h:] = 1.0 - r + + if cur_ov_w > 0: + r = get_ramp(cur_ov_w) + if x_idx > 0: + w_w[:cur_ov_w] = r + if x_end < w: + w_w[-cur_ov_w:] = 1.0 - r + + final_weight = w_h.view(1,1,1,-1,1) * w_w.view(1,1,1,1,-1) + + valid_d = min(tile_out.shape[2], result.shape[2]) + tile_out = tile_out[:, :, :valid_d, :, :] + + tile_out.mul_(final_weight) + + result[:, :, :valid_d, ys:ye, xs:xe] += tile_out + count[:, :, :, ys:ye, xs:xe] += final_weight + + del tile_out, final_weight, w_h, w_w + bar.update(1) + + result.div_(count.clamp(min=1e-6)) + + if result.device != x.device or result.dtype != x.dtype: + result = result.to(device=x.device, dtype=x.dtype) + + if x.shape[2] == 1 and sf_t == 1: + result = result.squeeze(2) + + return result + +_NORM_LIMIT = float("inf") +def get_norm_limit(): + return _NORM_LIMIT + + +def set_norm_limit(value: Optional[float] = None): + global _NORM_LIMIT + if value is None: + value = float("inf") + _NORM_LIMIT = value + +@contextmanager +def ignore_padding(model): + orig_padding = model.padding + model.padding = (0, 0, 0) + try: + yield + finally: + model.padding = orig_padding + +class MemoryState(Enum): + DISABLED = 0 + INITIALIZING = 1 + ACTIVE = 2 + UNSET = 3 + +def get_cache_size(conv_module, input_len, pad_len, dim=0): + dilated_kernel_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1 + output_len = (input_len + pad_len - dilated_kernel_size) // conv_module.stride[dim] + 1 + remain_len = ( + input_len + pad_len - ((output_len - 1) * conv_module.stride[dim] + dilated_kernel_size) + ) + overlap_len = dilated_kernel_size - conv_module.stride[dim] + cache_len = overlap_len + remain_len + + if output_len <= 0: + raise ValueError( + f"SeedVR2 VAE cache input is too short for convolution: input_len={input_len}, pad_len={pad_len}." + ) + return cache_len + +class DiagonalGaussianDistribution(object): + def __init__(self, parameters: torch.Tensor): + self.parameters = parameters + self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) + self.logvar = torch.clamp(self.logvar, BYTEDANCE_LOGVAR_CLAMP_MIN, BYTEDANCE_LOGVAR_CLAMP_MAX) + + def mode(self): + return self.mean + +class SpatialNorm(nn.Module): + def __init__( + self, + f_channels: int, + zq_channels: int, + ): + super().__init__() + self.norm_layer = ops.GroupNorm(num_channels=f_channels, num_groups=32, eps=1e-6, affine=True) + self.conv_y = ops.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) + self.conv_b = ops.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) + + def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor: + f_size = f.shape[-2:] + zq = F.interpolate(zq, size=f_size, mode="nearest") + norm_f = self.norm_layer(f) + new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) + return new_f + +class Attention(nn.Module): + def __init__( + self, + query_dim: int, + heads: int = 8, + dim_head: int = 64, + bias: bool = False, + norm_num_groups: Optional[int] = None, + spatial_norm_dim: Optional[int] = None, + out_bias: bool = True, + eps: float = 1e-5, + rescale_output_factor: float = 1.0, + residual_connection: bool = False, + ): + super().__init__() + + self.inner_dim = dim_head * heads + self.rescale_output_factor = rescale_output_factor + self.residual_connection = residual_connection + self.out_dim = query_dim + self.heads = heads + + if norm_num_groups is not None: + self.group_norm = ops.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True) + else: + self.group_norm = None + + if spatial_norm_dim is not None: + self.spatial_norm = SpatialNorm(f_channels=query_dim, zq_channels=spatial_norm_dim) + else: + self.spatial_norm = None + + self.to_q = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_k = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_v = ops.Linear(query_dim, self.inner_dim, bias=bias) + self.to_out = nn.ModuleList([]) + self.to_out.append(ops.Linear(self.inner_dim, self.out_dim, bias=out_bias)) + self.to_out.append(nn.Identity()) + + self.optimized_vae_attention = vae_attention() + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + + residual = hidden_states + if self.spatial_norm is not None: + hidden_states = self.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size = hidden_states.shape[0] + + if self.group_norm is not None: + hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = self.to_q(hidden_states) + key = self.to_k(hidden_states) + value = self.to_v(hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // self.heads + + query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2) + + if input_ndim == 4 and self.heads == 1: + query = query.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + key = key.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + value = value.squeeze(1).transpose(1, 2).reshape(batch_size, head_dim, height, width) + hidden_states = self.optimized_vae_attention(query, key, value).reshape(batch_size, self.heads, head_dim, height * width).transpose(2, 3) + else: + hidden_states = optimized_attention(query, key, value, heads = self.heads, skip_reshape=True, skip_output_reshape=True) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + hidden_states = self.to_out[0](hidden_states) + hidden_states = self.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if self.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / self.rescale_output_factor + + return hidden_states + + +def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor: + input_dtype = x.dtype + if isinstance(norm_layer, (nn.LayerNorm, nn.RMSNorm)): + if x.ndim == 4: + x = x.permute(0, 2, 3, 1) + x = norm_layer(x) + x = x.permute(0, 3, 1, 2) + return x.to(input_dtype) + if x.ndim == 5: + x = x.permute(0, 2, 3, 4, 1) + x = norm_layer(x) + x = x.permute(0, 4, 1, 2, 3) + return x.to(input_dtype) + if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)): + if x.ndim <= 4: + return norm_layer(x).to(input_dtype) + if x.ndim == 5: + b, c, t, h, w = x.shape + x = x.transpose(1, 2).reshape(b * t, c, h, w) + memory_occupy = x.numel() * x.element_size() / 1024**3 + if isinstance(norm_layer, nn.GroupNorm) and memory_occupy > get_norm_limit(): + num_chunks = min(BYTEDANCE_GN_CHUNKS_FP16 if x.element_size() == 2 else BYTEDANCE_GN_CHUNKS_FP32, norm_layer.num_groups) + if norm_layer.num_groups % num_chunks != 0: + raise ValueError( + f"SeedVR2 VAE GroupNorm groups must divide chunks: groups={norm_layer.num_groups}, chunks={num_chunks}." + ) + num_groups_per_chunk = norm_layer.num_groups // num_chunks + + weights = comfy.ops.cast_to_input(norm_layer.weight, x).chunk(num_chunks, dim=0) + biases = comfy.ops.cast_to_input(norm_layer.bias, x).chunk(num_chunks, dim=0) + x = list(x.chunk(num_chunks, dim=1)) + for i, (w, bias) in enumerate(zip(weights, biases)): + x[i] = F.group_norm(x[i], num_groups_per_chunk, w, bias, norm_layer.eps) + x[i] = x[i].to(input_dtype) + x = torch.cat(x, dim=1) + else: + x = norm_layer(x) + x = x.reshape((b, t, x.size(1), x.size(2), x.size(3))).transpose(1, 2) + return x.to(input_dtype) + raise TypeError(f"SeedVR2 VAE unsupported norm layer type: {type(norm_layer).__name__}") + +_receptive_field_t = Literal["half", "full"] + +def extend_head(tensor, times: int = 2, memory = None): + if memory is not None: + return torch.cat((memory.to(tensor), tensor), dim=2) + if times < 0: + raise ValueError(f"SeedVR2 VAE extend_head expected times >= 0, got {times}.") + if times == 0: + return tensor + else: + tile_repeat = [1] * tensor.ndim + tile_repeat[2] = times + return torch.cat(tensors=(torch.tile(tensor[:, :, :1], tile_repeat), tensor), dim=2) + +def cache_send_recv(tensor, cache_size, times, memory=None): + recv_buffer = None + + if memory is not None: + recv_buffer = memory.to(tensor[0]) + elif times > 0: + tile_repeat = [1] * tensor[0].ndim + tile_repeat[2] = times + recv_buffer = torch.tile(tensor[0][:, :, :1], tile_repeat) + + return recv_buffer + +class InflatedCausalConv3d(ops.Conv3d): + def __init__( + self, + *args, + inflation_mode, + **kwargs, + ): + self.inflation_mode = inflation_mode + super().__init__(*args, **kwargs) + self.temporal_padding = self.padding[0] + self.padding = (0, *self.padding[1:]) + self.memory_limit = float("inf") + self.logged_once = False + + def set_memory_limit(self, value: float): + self.memory_limit = value + + def _conv_forward(self, input, weight, bias, *args, **kwargs): + try: + return super()._conv_forward(input, weight, bias, *args, **kwargs) + except NotImplementedError: + # for: Could not run 'aten::cudnn_convolution' with arguments from the 'CPU' backend + if not self.logged_once: + logging.warning("VAE is on CPU for decoding. This is most likely due to not enough memory") + self.logged_once = True + return F.conv3d(input, weight, bias, *args, **kwargs) + + def memory_limit_conv( + self, + x, + *, + split_dim=3, + padding=(0, 0, 0, 0, 0, 0), + prev_cache=None, + ): + if math.isinf(self.memory_limit): + if prev_cache is not None: + x = torch.cat([prev_cache, x], dim=split_dim - 1) + return super().forward(x) + + shape = list(x.size()) + if prev_cache is not None: + shape[split_dim - 1] += prev_cache.size(split_dim - 1) + for i, pad_sum in enumerate((padding[4] + padding[5], padding[2] + padding[3], padding[0] + padding[1])): + shape[-3 + i] += pad_sum + memory_occupy = math.prod(shape) * x.element_size() / 1024**3 # GiB + if memory_occupy < self.memory_limit or split_dim == x.ndim: + x_concat = x + if prev_cache is not None: + x_concat = torch.cat([prev_cache, x], dim=split_dim - 1) + + def pad_and_forward(): + padded = F.pad(x_concat, padding, mode='constant', value=0.0) + if not padded.is_contiguous(): + padded = padded.contiguous() + with ignore_padding(self): + return torch.nn.Conv3d.forward(self, padded) + + return pad_and_forward() + + num_splits = math.ceil(memory_occupy / self.memory_limit) + size_per_split = x.size(split_dim) // num_splits + split_sizes = [size_per_split] * (num_splits - 1) + split_sizes += [x.size(split_dim) - sum(split_sizes)] + + x = list(x.split(split_sizes, dim=split_dim)) + if prev_cache is not None: + prev_cache = list(prev_cache.split(split_sizes, dim=split_dim)) + cache = None + for idx in range(len(x)): + if prev_cache is not None: + x[idx] = torch.cat([prev_cache[idx], x[idx]], dim=split_dim - 1) + + lpad_dim = (x[idx].ndim - split_dim - 1) * 2 + rpad_dim = lpad_dim + 1 + padding = list(padding) + padding[lpad_dim] = self.padding[split_dim - 2] if idx == 0 else 0 + padding[rpad_dim] = self.padding[split_dim - 2] if idx == len(x) - 1 else 0 + pad_len = padding[lpad_dim] + padding[rpad_dim] + padding = tuple(padding) + + next_cache = None + cache_len = cache.size(split_dim) if cache is not None else 0 + next_cache_size = get_cache_size( + conv_module=self, + input_len=x[idx].size(split_dim) + cache_len, + pad_len=pad_len, + dim=split_dim - 2, + ) + if next_cache_size != 0: + if next_cache_size > x[idx].size(split_dim): + raise ValueError( + f"SeedVR2 VAE cache size {next_cache_size} exceeds split size {x[idx].size(split_dim)}." + ) + next_cache = ( + x[idx].transpose(0, split_dim)[-next_cache_size:].transpose(0, split_dim) + ) + + x[idx] = self.memory_limit_conv( + x[idx], + split_dim=split_dim + 1, + padding=padding, + prev_cache=cache + ) + + cache = next_cache + + output = torch.cat(x, dim=split_dim) + return output + + def forward( + self, + input, + memory_state: MemoryState = MemoryState.UNSET, + memory_cache = None, + ) -> Tensor: + if memory_state == MemoryState.UNSET: + raise ValueError("SeedVR2 VAE convolution requires an explicit MemoryState.") + if memory_cache is None: + memory_cache = {} + if memory_state != MemoryState.ACTIVE: + memory_cache.pop(self, None) + if ( + math.isinf(self.memory_limit) + and torch.is_tensor(input) + ): + return self.basic_forward(input, memory_state, memory_cache) + return self.slicing_forward(input, memory_state, memory_cache) + + def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET, memory_cache = None): + mem_size = self.stride[0] - self.kernel_size[0] + memory = memory_cache.get(self) if memory_cache is not None else None + if (memory is not None) and (memory_state == MemoryState.ACTIVE): + input = extend_head(input, memory=memory, times=-1) + else: + input = extend_head(input, times=self.temporal_padding * 2) + next_memory = ( + input[:, :, mem_size:].detach() + if (mem_size != 0 and memory_state != MemoryState.DISABLED) + else None + ) + if memory_cache is not None and memory_state != MemoryState.DISABLED: + if next_memory is None: + memory_cache.pop(self, None) + else: + memory_cache[self] = next_memory + return super().forward(input) + + def slicing_forward( + self, + input, + memory_state: MemoryState = MemoryState.UNSET, + memory_cache = None, + ) -> Tensor: + if memory_cache is None: + memory_cache = {} + squeeze_out = False + if torch.is_tensor(input): + input = [input] + squeeze_out = True + + cache_size = self.kernel_size[0] - self.stride[0] + memory = memory_cache.get(self) if memory_cache is not None else None + cache = cache_send_recv( + input, cache_size=cache_size, memory=memory, times=self.temporal_padding * 2 + ) + + if ( + memory_state in [MemoryState.INITIALIZING, MemoryState.ACTIVE] + and cache_size != 0 + ): + if cache_size > input[-1].size(2) and cache is not None and len(input) == 1: + input[0] = torch.cat([cache, input[0]], dim=2) + cache = None + if cache_size <= input[-1].size(2): + memory_cache[self] = input[-1][:, :, -cache_size:].detach().contiguous() + + padding = tuple(x for x in reversed(self.padding) for _ in range(2)) + for i in range(len(input)): + next_cache = None + cache_size = 0 + if i < len(input) - 1: + cache_len = cache.size(2) if cache is not None else 0 + cache_size = get_cache_size(self, input[i].size(2) + cache_len, pad_len=0) + if cache_size != 0: + if cache_size > input[i].size(2) and cache is not None: + input[i] = torch.cat([cache, input[i]], dim=2) + cache = None + if cache_size > input[i].size(2): + raise ValueError(f"SeedVR2 VAE cache size {cache_size} exceeds input length {input[i].size(2)}.") + next_cache = input[i][:, :, -cache_size:] + + input[i] = self.memory_limit_conv( + input[i], + padding=padding, + prev_cache=cache + ) + + cache = next_cache + + return input[0] if squeeze_out else input + +def remove_head(tensor: Tensor, times: int = 1) -> Tensor: + if times == 0: + return tensor + return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2) + +class Upsample3D(nn.Module): + + def __init__( + self, + channels, + out_channels = None, + inflation_mode = "tail", + temporal_up: bool = False, + spatial_up: bool = True, + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + + conv = InflatedCausalConv3d( + self.channels, + self.out_channels, + 3, + padding=1, + inflation_mode=inflation_mode, + ) + + self.temporal_up = temporal_up + self.spatial_up = spatial_up + self.temporal_ratio = 2 if temporal_up else 1 + self.spatial_ratio = 2 if spatial_up else 1 + + upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio + self.upscale_conv = ops.Conv3d( + self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0 + ) + + self.conv = conv + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state=None, + memory_cache=None, + ) -> torch.FloatTensor: + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 upsample expected {self.channels} channels, got {hidden_states.shape[1]}.") + + hidden_states = self.upscale_conv(hidden_states) + b, channels, f, h, w = hidden_states.shape + c = channels // (self.spatial_ratio * self.spatial_ratio * self.temporal_ratio) + hidden_states = hidden_states.view(b, self.spatial_ratio, self.spatial_ratio, self.temporal_ratio, c, f, h, w) + hidden_states = hidden_states.permute(0, 4, 5, 3, 6, 1, 7, 2).reshape( + b, + c, + f * self.temporal_ratio, + h * self.spatial_ratio, + w * self.spatial_ratio, + ) + + if self.temporal_up and memory_state != MemoryState.ACTIVE: + hidden_states = remove_head(hidden_states) + + hidden_states = self.conv(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class Downsample3D(nn.Module): + def __init__( + self, + channels, + out_channels = None, + inflation_mode = "tail", + spatial_down: bool = False, + temporal_down: bool = False, + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.temporal_down = temporal_down + self.spatial_down = spatial_down + + self.temporal_ratio = 2 if temporal_down else 1 + self.spatial_ratio = 2 if spatial_down else 1 + + self.temporal_kernel = 3 if temporal_down else 1 + self.spatial_kernel = 3 if spatial_down else 1 + + self.conv = InflatedCausalConv3d( + self.channels, + self.out_channels, + kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel), + stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), + padding=(1 if self.temporal_down else 0, 0, 0), + inflation_mode=inflation_mode, + ) + + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 downsample expected {self.channels} channels, got {hidden_states.shape[1]}.") + + if self.spatial_down: + pad = (0, 1, 0, 1) + hidden_states = F.pad(hidden_states, pad, mode="constant", value=0) + + if hidden_states.shape[1] != self.channels: + raise ValueError(f"SeedVR2 downsample expected {self.channels} channels after padding, got {hidden_states.shape[1]}.") + + hidden_states = self.conv(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class ResnetBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: Optional[int] = None, + temb_channels: int = 512, + groups: int = 32, + groups_out: Optional[int] = None, + eps: float = 1e-6, + output_scale_factor: float = 1.0, + skip_time_act: bool = False, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = in_channels if out_channels is None else out_channels + self.output_scale_factor = output_scale_factor + self.skip_time_act = skip_time_act + self.nonlinearity = nn.SiLU() + if temb_channels is not None: + self.time_emb_proj = ops.Linear(temb_channels, self.out_channels) + else: + self.time_emb_proj = None + self.norm1 = ops.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True) + if groups_out is None: + groups_out = groups + self.norm2 = ops.GroupNorm(num_groups=groups_out, num_channels=self.out_channels, eps=eps, affine=True) + self.use_in_shortcut = self.in_channels != self.out_channels + self.conv1 = InflatedCausalConv3d( + self.in_channels, + self.out_channels, + kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3), + stride=1, + padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1), + inflation_mode=inflation_mode, + ) + + self.conv2 = InflatedCausalConv3d( + self.out_channels, + self.out_channels, + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.conv_shortcut = None + if self.use_in_shortcut: + self.conv_shortcut = InflatedCausalConv3d( + self.in_channels, + self.out_channels, + kernel_size=1, + stride=1, + padding=0, + bias=True, + inflation_mode=inflation_mode, + ) + + def forward(self, input_tensor, temb, memory_state = None, memory_cache = None): + hidden_states = input_tensor + + hidden_states = causal_norm_wrapper(self.norm1, hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + + hidden_states = self.conv1(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + if self.time_emb_proj is not None: + if not self.skip_time_act: + temb = self.nonlinearity(temb) + temb = self.time_emb_proj(temb)[:, :, None, None] + + if temb is not None: + hidden_states = hidden_states + temb + + hidden_states = causal_norm_wrapper(self.norm2, hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + + hidden_states = self.conv2(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + if self.conv_shortcut is not None: + input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state, memory_cache=memory_cache) + + output_tensor = (input_tensor + hidden_states) / self.output_scale_factor + + return output_tensor + + +class DownEncoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_groups: int = 32, + output_scale_factor: float = 1.0, + add_downsample: bool = True, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_down: bool = True, + spatial_down: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=None, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample3D( + out_channels, + out_channels=out_channels, + temporal_down=temporal_down, + spatial_down=spatial_down, + inflation_mode=inflation_mode, + ) + ] + ) + else: + self.downsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, memory_cache=memory_cache) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class UpDecoderBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_groups: int = 32, + output_scale_factor: float = 1.0, + add_upsample: bool = True, + temb_channels: Optional[int] = None, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up: bool = True, + spatial_up: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + + resnets.append( + ResnetBlock3D( + in_channels=input_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_upsample: + self.upsamplers = nn.ModuleList( + [ + Upsample3D( + out_channels, + out_channels=out_channels, + temporal_up=temporal_up, + spatial_up=spatial_up, + inflation_mode=inflation_mode, + ) + ] + ) + else: + self.upsamplers = None + + def forward( + self, + hidden_states: torch.FloatTensor, + temb: Optional[torch.FloatTensor] = None, + memory_state=None, + memory_cache=None, + ) -> torch.FloatTensor: + for resnet in self.resnets: + hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, memory_cache=memory_cache) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class UNetMidBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + temb_channels: int, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", # default, spatial + resnet_groups: int = 32, + add_attention: bool = True, + attention_head_dim: int = 1, + output_scale_factor: float = 1.0, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) + self.add_attention = add_attention + + resnets = [ + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ] + attentions = [] + + if attention_head_dim is None: + attention_head_dim = in_channels + + for _ in range(num_layers): + if self.add_attention: + attentions.append( + Attention( + in_channels, + heads=in_channels // attention_head_dim, + dim_head=attention_head_dim, + rescale_output_factor=output_scale_factor, + eps=resnet_eps, + norm_num_groups=( + resnet_groups if resnet_time_scale_shift == "default" else None + ), + spatial_norm_dim=( + temb_channels if resnet_time_scale_shift == "spatial" else None + ), + residual_connection=True, + bias=True, + ) + ) + else: + attentions.append(None) + + resnets.append( + ResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + output_scale_factor=output_scale_factor, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + def forward(self, hidden_states, temb=None, memory_state=None, memory_cache=None): + video_length = hidden_states.size(2) + hidden_states = self.resnets[0](hidden_states, temb, memory_state=memory_state, memory_cache=memory_cache) + for attn, resnet in zip(self.attentions, self.resnets[1:]): + if attn is not None: + b, c, f, h, w = hidden_states.shape + hidden_states = hidden_states.transpose(1, 2).reshape(b * f, c, h, w) + hidden_states = attn(hidden_states, temb=temb) + hidden_states = hidden_states.reshape(b, video_length, c, h, w).transpose(1, 2) + hidden_states = resnet(hidden_states, temb, memory_state=memory_state, memory_cache=memory_cache) + + return hidden_states + + +class Encoder3D(nn.Module): + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + down_block_types: Tuple[str, ...] = ("DownEncoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + mid_block_add_attention=True, + temporal_down_num: int = 2, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_down_num = temporal_down_num + + self.conv_in = InflatedCausalConv3d( + in_channels, + block_out_channels[0], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.down_blocks = nn.ModuleList([]) + + output_channel = block_out_channels[0] + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1 + + if down_block_type != "DownEncoderBlock3D": + raise ValueError(f"SeedVR2 encoder only supports DownEncoderBlock3D, got {down_block_type}.") + + down_block = DownEncoderBlock3D( + num_layers=self.layers_per_block, + in_channels=input_channel, + out_channels=output_channel, + add_downsample=not is_final_block, + resnet_eps=1e-6, + resnet_groups=norm_num_groups, + temporal_down=is_temporal_down_block, + spatial_down=True, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.down_blocks.append(down_block) + + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + output_scale_factor=1, + resnet_time_scale_shift="default", + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=None, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.conv_norm_out = ops.GroupNorm( + num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + + conv_out_channels = 2 * out_channels + self.conv_out = InflatedCausalConv3d( + block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + + def forward( + self, + sample: torch.FloatTensor, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + sample = sample.to(next(self.parameters()).device) + sample = self.conv_in(sample, memory_state=memory_state, memory_cache=memory_cache) + for down_block in self.down_blocks: + sample = down_block(sample, memory_state=memory_state, memory_cache=memory_cache) + + sample = self.mid_block(sample, memory_state=memory_state, memory_cache=memory_cache) + + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state, memory_cache=memory_cache) + + return sample + + +class Decoder3D(nn.Module): + + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + up_block_types: Tuple[str, ...] = ("UpDecoderBlock3D",), + block_out_channels: Tuple[int, ...] = (64,), + layers_per_block: int = 2, + norm_num_groups: int = 32, + mid_block_add_attention=True, + inflation_mode = "tail", + time_receptive_field: _receptive_field_t = "half", + temporal_up_num: int = 2, + ): + super().__init__() + self.layers_per_block = layers_per_block + self.temporal_up_num = temporal_up_num + + self.conv_in = InflatedCausalConv3d( + in_channels, + block_out_channels[-1], + kernel_size=3, + stride=1, + padding=1, + inflation_mode=inflation_mode, + ) + + self.mid_block = None + self.up_blocks = nn.ModuleList([]) + + temb_channels = None + + self.mid_block = UNetMidBlock3D( + in_channels=block_out_channels[-1], + resnet_eps=1e-6, + output_scale_factor=1, + resnet_time_scale_shift="default", + attention_head_dim=block_out_channels[-1], + resnet_groups=norm_num_groups, + temb_channels=temb_channels, + add_attention=mid_block_add_attention, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + reversed_block_out_channels = list(reversed(block_out_channels)) + output_channel = reversed_block_out_channels[0] + for i, up_block_type in enumerate(up_block_types): + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + + is_final_block = i == len(block_out_channels) - 1 + is_temporal_up_block = i < self.temporal_up_num + if up_block_type != "UpDecoderBlock3D": + raise ValueError(f"SeedVR2 decoder only supports UpDecoderBlock3D, got {up_block_type}.") + up_block = UpDecoderBlock3D( + num_layers=self.layers_per_block + 1, + in_channels=prev_output_channel, + out_channels=output_channel, + add_upsample=not is_final_block, + resnet_eps=1e-6, + resnet_groups=norm_num_groups, + temb_channels=temb_channels, + temporal_up=is_temporal_up_block, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + self.up_blocks.append(up_block) + prev_output_channel = output_channel + + self.conv_norm_out = ops.GroupNorm( + num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6 + ) + self.conv_act = nn.SiLU() + self.conv_out = InflatedCausalConv3d( + block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode + ) + + + def forward( + self, + sample: torch.FloatTensor, + latent_embeds: Optional[torch.FloatTensor] = None, + memory_state = None, + memory_cache = None, + ) -> torch.FloatTensor: + + sample = sample.to(next(self.parameters()).device) + sample = self.conv_in(sample, memory_state=memory_state, memory_cache=memory_cache) + + upscale_dtype = next(iter(self.up_blocks.parameters())).dtype + sample = self.mid_block(sample, latent_embeds, memory_state=memory_state, memory_cache=memory_cache) + sample = sample.to(upscale_dtype) + + for up_block in self.up_blocks: + sample = up_block(sample, latent_embeds, memory_state=memory_state, memory_cache=memory_cache) + + sample = causal_norm_wrapper(self.conv_norm_out, sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample, memory_state=memory_state, memory_cache=memory_cache) + + return sample + +class VideoAutoencoderKL(nn.Module): + def __init__( + self, + in_channels: int = 3, + out_channels: int = 3, + layers_per_block: int = 2, + latent_channels: int = SEEDVR2_LATENT_CHANNELS, + norm_num_groups: int = 32, + temporal_scale_num: int = 2, + inflation_mode = "pad", + time_receptive_field: _receptive_field_t = "full", + slicing_sample_min_size = BYTEDANCE_SLICING_SAMPLE_MIN, + ): + self.slicing_sample_min_size = slicing_sample_min_size + self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num) + block_out_channels = BYTEDANCE_BLOCK_OUT_CHANNELS + down_block_types = ("DownEncoderBlock3D",) * 4 + up_block_types = ("UpDecoderBlock3D",) * 4 + super().__init__() + + self.encoder = Encoder3D( + in_channels=in_channels, + out_channels=latent_channels, + down_block_types=down_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + norm_num_groups=norm_num_groups, + temporal_down_num=temporal_scale_num, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.decoder = Decoder3D( + in_channels=latent_channels, + out_channels=out_channels, + up_block_types=up_block_types, + block_out_channels=block_out_channels, + layers_per_block=layers_per_block, + norm_num_groups=norm_num_groups, + temporal_up_num=temporal_scale_num, + inflation_mode=inflation_mode, + time_receptive_field=time_receptive_field, + ) + + self.use_slicing = True + + def encode(self, x: torch.FloatTensor, return_dict: bool = True): + h = self.slicing_encode(x) + posterior = DiagonalGaussianDistribution(h).mode() + + if not return_dict: + return (posterior,) + + return posterior + + def decode_( + self, z: torch.Tensor, return_dict: bool = True + ): + decoded = self.slicing_decode(z) + + if not return_dict: + return (decoded,) + + return decoded + + def _encode( + self, x, memory_state = MemoryState.DISABLED, memory_cache = None + ) -> torch.Tensor: + _x = x.to(self.device) + h = self.encoder(_x, memory_state=memory_state, memory_cache=memory_cache) + return h.to(x.device) + + def _decode( + self, z, memory_state = MemoryState.DISABLED, memory_cache = None + ) -> torch.Tensor: + _z = z.to(self.device) + output = self.decoder(_z, memory_state=memory_state, memory_cache=memory_cache) + return output.to(z.device) + + def slicing_encode(self, x: torch.Tensor) -> torch.Tensor: + if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size: + memory_cache = {} + split_size = max( + self.slicing_sample_min_size, + getattr(self, "temporal_downsample_factor", 1), + ) + x_slices = list(x[:, :, 1:].split(split_size=split_size, dim=2)) + min_active_len = getattr(self, "temporal_downsample_factor", 1) + if len(x_slices) > 1 and x_slices[-1].shape[2] < min_active_len: + x_slices[-2] = torch.cat((x_slices[-2], x_slices[-1]), dim=2) + x_slices.pop() + encoded_slices = [ + self._encode( + torch.cat((x[:, :, :1], x_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + memory_cache=memory_cache, + ) + ] + for x_idx in range(1, len(x_slices)): + encoded_slices.append( + self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE, memory_cache=memory_cache) + ) + out = torch.cat(encoded_slices, dim=2) + return out + else: + return self._encode(x) + + def slicing_decode(self, z: torch.Tensor) -> torch.Tensor: + if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size: + memory_cache = {} + z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size, dim=2) + decoded_slices = [ + self._decode( + torch.cat((z[:, :, :1], z_slices[0]), dim=2), + memory_state=MemoryState.INITIALIZING, + memory_cache=memory_cache, + ) + ] + for z_idx in range(1, len(z_slices)): + decoded_slices.append( + self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE, memory_cache=memory_cache) + ) + out = torch.cat(decoded_slices, dim=2) + return out + else: + return self._decode(z) + + def forward(self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all"): + def _unwrap(value): + return value[0] if isinstance(value, tuple) else value + + if mode == "encode": + return _unwrap(self.encode(x)) + if mode == "decode": + return _unwrap(self.decode_(x)) + if mode == "all": + latent = _unwrap(self.encode(x)) + return _unwrap(self.decode_(latent)) + raise ValueError(f"Unknown SeedVR2 VAE forward mode: {mode}") + +class VideoAutoencoderKLWrapper(VideoAutoencoderKL): + def __init__( + self, + spatial_downsample_factor = 8, + temporal_downsample_factor = 4, + ): + self.spatial_downsample_factor = spatial_downsample_factor + self.temporal_downsample_factor = temporal_downsample_factor + super().__init__() + self.set_memory_limit(BYTEDANCE_VAE_CONV_MEM_GIB, BYTEDANCE_VAE_NORM_MEM_GIB) + + def forward(self, x: torch.FloatTensor): + z, p = self._encode_with_raw_latent(x) + x = self.decode(z) + return x, z, p + + def _encode_with_raw_latent(self, x): + if x.ndim == 4: + x = x.unsqueeze(2) + self.device = x.device + p = super().encode(x) + z = p.squeeze(2) + return z, p + + def encode(self, x): + z, _ = self._encode_with_raw_latent(x) + return z + + def decode(self, z, seedvr2_tiling=None): + seedvr2_tiling = {} if seedvr2_tiling is None else seedvr2_tiling + if not isinstance(seedvr2_tiling, dict): + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: `seedvr2_tiling` must be a dict; " + f"got {type(seedvr2_tiling).__name__} with value {seedvr2_tiling!r}." + ) + + if z.ndim == 5: + _, c, _, _, _ = z.shape + if c != SEEDVR2_LATENT_CHANNELS: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: 5-D latent input must " + f"have {SEEDVR2_LATENT_CHANNELS} channels; got shape {tuple(z.shape)}." + ) + latent = z + elif z.ndim == 4: + b, tc, h, w = z.shape + if tc % SEEDVR2_LATENT_CHANNELS != 0: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: 4-D latent input must " + f"use collapsed channel layout (B, {SEEDVR2_LATENT_CHANNELS}*T, H, W); " + f"got shape {tuple(z.shape)}." + ) + latent = z.reshape(b, SEEDVR2_LATENT_CHANNELS, -1, h, w) + else: + raise RuntimeError( + "SeedVR2 VideoAutoencoderKLWrapper.decode: latent input must be " + f"4-D collapsed (B, {SEEDVR2_LATENT_CHANNELS}*T, H, W) or " + f"5-D (B, {SEEDVR2_LATENT_CHANNELS}, T, H, W); " + f"got shape {tuple(z.shape)}." + ) + scale = BYTEDANCE_VAE_SCALING_FACTOR + shift = BYTEDANCE_VAE_SHIFTING_FACTOR + latent = latent / scale + shift + + self.device = latent.device + enable_tiling = seedvr2_tiling.get("enable_tiling", False) + + if enable_tiling: + decode_seedvr2_args = dict(seedvr2_tiling) + decode_seedvr2_args.pop("enable_tiling", None) + tile_h, tile_w = decode_seedvr2_args.get("tile_size", (512, 512)) + ov_h, ov_w = decode_seedvr2_args.get("tile_overlap", (64, 64)) + decode_seedvr2_args["tile_overlap"] = ( + min(ov_h, max(0, tile_h - 8)), + min(ov_w, max(0, tile_w - 8)), + ) + x = tiled_vae(latent, self, **decode_seedvr2_args, encode=False) + if x.ndim == 4: + # tiled_vae squeezes the temporal axis when + # temporal_downsample_factor == 1 AND latent T == 1 + # (see tiled_vae line 179-180); re-add it so the post-decode + # pipeline can keep batch and time distinct on the tiled path. + x = x.unsqueeze(2) + else: + x = super().decode_(latent) + + h, w = x.shape[-2:] + w2 = w - (w % 2) + h2 = h - (h % 2) + x = x[..., :h2, :w2] + + return x + + def decode_tiled(self, z, tile_x=32, tile_y=32, overlap=8, tile_t=None, overlap_t=None): + # SeedVR2's causal VAE owns temporal via the MemoryState cache; external + # temporal tiling breaks that continuity, so only spatial tiling is applied. + sf = self.spatial_downsample_factor + seedvr2_tiling = { + "enable_tiling": True, + "tile_size": (tile_y * sf, tile_x * sf), + "tile_overlap": (overlap * sf, overlap * sf), + "temporal_size": None, + "temporal_overlap": None, + } + return self.decode(z, seedvr2_tiling=seedvr2_tiling) + + def encode_tiled(self, x, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + # External temporal tiling knobs are discarded; the causal VAE keeps its + # own internal MemoryState slicing. + if tile_y is None: + tile_y = 512 + if tile_x is None: + tile_x = 512 + if overlap is None: + overlap_y = 64 + overlap_x = 64 + else: + overlap_y = overlap + overlap_x = overlap + overlap_y = min(overlap_y, max(0, tile_y - 8)) + overlap_x = min(overlap_x, max(0, tile_x - 8)) + self.device = x.device + return tiled_vae( + x, + self, + tile_size=(tile_y, tile_x), + tile_overlap=(overlap_y, overlap_x), + temporal_size=None, + temporal_overlap=None, + encode=True, + ) + + def comfy_format_encoded(self, samples): + if samples.ndim == 4: + samples = samples.unsqueeze(2) + samples = samples.contiguous() + samples = samples * BYTEDANCE_VAE_SCALING_FACTOR + return samples + + def comfy_memory_used_decode(self, shape): + bytes_per_output_pixel = 160 + + def output_pixels(latent_t, latent_h, latent_w): + output_t = max(1, (latent_t - 1) * 4 + 1) + return output_t * latent_h * 8 * latent_w * 8 + + # SeedVR2 decode performs full-frame LAB histogram matching: fp32 channels + # plus int64 sort indices dominate peak memory, not the VAE weight dtype. + if len(shape) == 5: + candidates = [] + if shape[1] == SEEDVR2_LATENT_CHANNELS: + candidates.append((shape[2], shape[3], shape[4])) + if shape[-1] == SEEDVR2_LATENT_CHANNELS: + candidates.append((shape[1], shape[2], shape[3])) + if len(candidates) == 0: + candidates.append((shape[2], shape[3], shape[4])) + pixels = max(output_pixels(*candidate) for candidate in candidates) + elif len(shape) == 4: + latent_t = max(1, (shape[1] + SEEDVR2_LATENT_CHANNELS - 1) // SEEDVR2_LATENT_CHANNELS) + pixels = output_pixels(latent_t, shape[2], shape[3]) + else: + pixels = output_pixels(1, shape[-2], shape[-1]) + return pixels * bytes_per_output_pixel + + def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]): + set_norm_limit(norm_max_mem) + for m in self.modules(): + if isinstance(m, InflatedCausalConv3d): + m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf")) diff --git a/comfy/ldm/wan/model.py b/comfy/ldm/wan/model.py index 1c9782a38..c042e93c4 100644 --- a/comfy/ldm/wan/model.py +++ b/comfy/ldm/wan/model.py @@ -552,6 +552,7 @@ class WanModel(torch.nn.Module): List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] """ # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] transformer_options["grid_sizes"] = grid_sizes @@ -564,11 +565,13 @@ class WanModel(torch.nn.Module): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # In-context reference (Bernini) context_latents = kwargs.get("context_latents", None) @@ -589,6 +592,7 @@ class WanModel(torch.nn.Module): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -604,6 +608,11 @@ class WanModel(torch.nn.Module): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -777,6 +786,7 @@ class VaceWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] transformer_options["grid_sizes"] = grid_sizes @@ -807,6 +817,7 @@ class VaceWanModel(WanModel): x_orig = x patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -822,6 +833,11 @@ class VaceWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + ii = self.vace_layers_mapping.get(i, None) if ii is not None: for iii in range(len(c)): @@ -887,6 +903,7 @@ class CameraWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) if self.control_adapter is not None and camera_conditions is not None: x = x + self.control_adapter(camera_conditions).to(x.dtype) @@ -909,6 +926,7 @@ class CameraWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -924,6 +942,11 @@ class CameraWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -1335,6 +1358,7 @@ class WanModel_S2V(WanModel): # embeddings bs, _, time, height, width = x.shape + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) if control_video is not None: x = x + self.cond_encoder(control_video) @@ -1379,6 +1403,7 @@ class WanModel_S2V(WanModel): context = self.text_embedding(context) patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1393,6 +1418,12 @@ class WanModel_S2V(WanModel): x = out["img"] else: x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + if audio_emb is not None: x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len) # head @@ -1599,6 +1630,7 @@ class HumoWanModel(WanModel): bs, _, time, height, width = x.shape # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) grid_sizes = x.shape[2:] x = x.flatten(2).transpose(1, 2) @@ -1630,6 +1662,7 @@ class HumoWanModel(WanModel): audio = None patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1645,6 +1678,11 @@ class HumoWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) @@ -1660,8 +1698,14 @@ class SCAILWanModel(WanModel): def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs): + x_input = x + + img_offset = 0 if reference_latent is not None: x = torch.cat((reference_latent, x), dim=2) + img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \ + (reference_latent.shape[3] // self.patch_size[1]) * \ + (reference_latent.shape[4] // self.patch_size[2]) # embeddings x = self.patch_embedding(x.float()).to(x.dtype) @@ -1697,6 +1741,7 @@ class SCAILWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -1712,6 +1757,11 @@ class SCAILWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + # head x = self.head(x, e) diff --git a/comfy/ldm/wan/model_animate.py b/comfy/ldm/wan/model_animate.py index 84d7adec4..9ebe5694b 100644 --- a/comfy/ldm/wan/model_animate.py +++ b/comfy/ldm/wan/model_animate.py @@ -493,6 +493,7 @@ class AnimateWanModel(WanModel): **kwargs, ): # embeddings + x_input = x x = self.patch_embedding(x.float()).to(x.dtype) x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values) grid_sizes = x.shape[2:] @@ -505,11 +506,13 @@ class AnimateWanModel(WanModel): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # context context = self.text_embedding(context) @@ -522,6 +525,7 @@ class AnimateWanModel(WanModel): context_img_len = clip_fea.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -537,6 +541,11 @@ class AnimateWanModel(WanModel): else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + if i % 5 == 0 and motion_vec is not None: x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec) diff --git a/comfy/ldm/wan/model_wandancer.py b/comfy/ldm/wan/model_wandancer.py index 3caef6dc5..aeec1d725 100644 --- a/comfy/ldm/wan/model_wandancer.py +++ b/comfy/ldm/wan/model_wandancer.py @@ -111,6 +111,7 @@ class WanDancerModel(WanModel): def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs): # embeddings + x_input = x if int(fps + 0.5) != 30: x = self.patch_embedding_global(x.float()).to(x.dtype) else: @@ -128,11 +129,13 @@ class WanDancerModel(WanModel): e0 = self.time_projection(e).unflatten(2, (6, self.dim)) full_ref = None + img_offset = 0 if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline full_ref = kwargs.get("reference_latent", None) if full_ref is not None: full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2) x = torch.concat((full_ref, x), dim=1) + img_offset = full_ref.shape[1] # context context = self.text_embedding(context) @@ -163,6 +166,7 @@ class WanDancerModel(WanModel): context_img_len += clip_fea_ref.shape[-2] patches_replace = transformer_options.get("patches_replace", {}) + patches = transformer_options.get("patches", {}) blocks_replace = patches_replace.get("dit", {}) transformer_options["total_blocks"] = len(self.blocks) transformer_options["block_type"] = "double" @@ -177,6 +181,12 @@ class WanDancerModel(WanModel): x = out["img"] else: x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options) + + if "double_block" in patches: + for p in patches["double_block"]: + out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options}) + x = out["img"] + if audio_emb is not None: x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale) diff --git a/comfy/ldm/wan/uni3c.py b/comfy/ldm/wan/uni3c.py new file mode 100644 index 000000000..827ad2339 --- /dev/null +++ b/comfy/ldm/wan/uni3c.py @@ -0,0 +1,149 @@ +# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C +# Converted from the original diffusers based implementation. +import torch +import torch.nn as nn + +from comfy.ldm.flux.layers import EmbedND +from .model import WanSelfAttention + + +class Uni3CLayerNormZero(nn.Module): + def __init__( + self, + conditioning_dim, + embedding_dim, + eps=1e-5, + device=None, dtype=None, operations=None + ): + super().__init__() + self.silu = nn.SiLU() + self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype) + self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype) + + def forward(self, x, temb): + shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1) + x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] + return x, gate[:, None, :] + + +class Uni3CAttentionBlock(nn.Module): + def __init__( + self, + dim, + ffn_dim, + num_heads, + time_embed_dim=5120, + eps=1e-6, + device=None, dtype=None, operations=None + ): + super().__init__() + operation_settings = {"operations": operations, "device": device, "dtype": dtype} + self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) + self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings) + self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations) + self.ffn = nn.Sequential( + operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'), + operations.Linear(ffn_dim, dim, device=device, dtype=dtype)) + + def forward(self, x, temb, freqs): + norm_x, gate_msa = self.norm1(x, temb) + x = x + gate_msa * self.self_attn(norm_x, freqs) + norm_x, gate_ff = self.norm2(x, temb) + x = x + gate_ff * self.ffn(norm_x) + return x + + +class MaskCamEmbed(nn.Module): + def __init__( + self, + add_channels=7, + mid_channels=256, + conv_out_dim=5120, + device=None, dtype=None, operations=None + ): + super().__init__() + self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning + self.mask_proj = nn.Sequential( + operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype), + operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype), + nn.SiLU()) + self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype) + + def forward(self, add_inputs): + add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0) + add_embeds = self.mask_proj(add_padded) + add_embeds = self.mask_zero_proj(add_embeds) + add_embeds = add_embeds.flatten(2).transpose(1, 2) + return add_embeds + + +class WanUni3CControlnet(nn.Module): + def __init__( + self, + in_channels=36, + conv_out_dim=5120, + dim=1024, + ffn_dim=8192, + num_heads=16, + num_layers=20, + time_embed_dim=5120, + out_proj_dim=5120, + add_channels=7, + mid_channels=256, + device=None, dtype=None, operations=None + ): + super().__init__() + patch_size = (1, 2, 2) + self.num_layers = num_layers + + self.controlnet_patch_embedding = operations.Conv3d( + in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32) + self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations) + + if conv_out_dim != dim: + self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype) + else: + self.proj_in = nn.Identity() + + self.controlnet_blocks = nn.ModuleList([ + Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations) + for _ in range(num_layers)]) + self.proj_out = nn.ModuleList([ + operations.Linear(dim, out_proj_dim, device=device, dtype=dtype) + for _ in range(num_layers)]) + + head_dim = dim // num_heads + self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)]) + + def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None): + img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype) + img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1) + img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1) + img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1) + img_ids = img_ids.reshape(1, -1, img_ids.shape[-1]) + freqs = self.rope_embedder(img_ids).movedim(1, 2) + return freqs + + def process_input(self, control_input, render_mask=None, camera_embedding=None): + # render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet + hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype) + t_len, h_len, w_len = hidden.shape[2:] + freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype) + hidden = hidden.flatten(2).transpose(1, 2) + + add_inputs = None + if camera_embedding is not None and render_mask is not None: + add_inputs = torch.cat([render_mask, camera_embedding], dim=1) + elif render_mask is not None: + add_inputs = render_mask + + if add_inputs is not None: + hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype)) + + hidden = self.proj_in(hidden) + return hidden, freqs + + def forward_block(self, block_index, hidden, temb, freqs): + hidden = self.controlnet_blocks[block_index](hidden, temb, freqs) + residual = self.proj_out[block_index](hidden) + return hidden, residual diff --git a/comfy/lora.py b/comfy/lora.py index 2c8d0f0bf..427cf98aa 100644 --- a/comfy/lora.py +++ b/comfy/lora.py @@ -326,6 +326,17 @@ def model_lora_keys_unet(model, key_map={}): key_map["transformer.{}".format(key_lora)] = k key_map["lycoris_{}".format(key_lora.replace(".", "_"))] = k #SimpleTuner lycoris format + if isinstance(model, comfy.model_base.Krea2): + diffusers_keys = comfy.utils.krea2_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") + for k in diffusers_keys: + if k.endswith(".weight"): + to = diffusers_keys[k] + key_lora = k[:-len(".weight")] + key_map["diffusion_model.{}".format(key_lora)] = to + key_map["transformer.{}".format(key_lora)] = to + key_map["lycoris_{}".format(key_lora.replace(".", "_"))] = to + key_map[key_lora] = to + if isinstance(model, comfy.model_base.Lumina2): diffusers_keys = comfy.utils.z_image_to_diffusers(model.model_config.unet_config, output_prefix="diffusion_model.") for k in diffusers_keys: diff --git a/comfy/model_base.py b/comfy/model_base.py index 264dbb9b3..6631c9eb0 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -21,6 +21,7 @@ import comfy.ldm.hunyuan3dv2_1.hunyuandit import torch import logging import comfy.ldm.lightricks.av_model +import comfy.ldm.minimax.model import comfy.ldm.lightricks.symmetric_patchifier import comfy.context_windows from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep @@ -55,9 +56,13 @@ import comfy.ldm.pixeldit.model import comfy.ldm.pixeldit.pid import comfy.ldm.ace.model import comfy.ldm.omnigen.omnigen2 +import comfy.ldm.seedvr.model import comfy.ldm.boogu.model import comfy.ldm.qwen_image.model +import comfy.ldm.mage_flow.model +import comfy.ldm.joyimage.model import comfy.ldm.ideogram4.model +import comfy.ldm.krea2.model import comfy.ldm.kandinsky5.model import comfy.ldm.anima.model import comfy.ldm.ace.ace_step15 @@ -165,6 +170,7 @@ class BaseModel(torch.nn.Module): else: operations = model_config.custom_operations self.diffusion_model = unet_model(**unet_config, device=device, operations=operations) + self.diffusion_model.requires_grad_(False) self.diffusion_model.eval() if comfy.model_management.force_channels_last(): self.diffusion_model.to(memory_format=torch.channels_last) @@ -931,6 +937,17 @@ class HunyuanDiT(BaseModel): out['image_meta_size'] = comfy.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]])) return out +class SeedVR2(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.seedvr.model.NaDiT) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + condition = kwargs.get("condition", None) + if condition is not None: + out["condition"] = comfy.conds.CONDRegular(condition) + return out + class PixArt(BaseModel): def __init__(self, model_config, model_type=ModelType.EPS, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.pixart.pixartms.PixArtMS) @@ -2010,11 +2027,11 @@ class WAN22_WanDancer(WAN21): fps = kwargs.get("fps", None) if fps is not None: - out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps])) + out['fps'] = comfy.conds.CONDConstant(fps) audio_inject_scale = kwargs.get("audio_inject_scale", None) if audio_inject_scale is not None: - out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale])) + out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale) return out class Hunyuan3Dv2(BaseModel): @@ -2047,6 +2064,57 @@ class Hunyuan3Dv2_1(BaseModel): out['guidance'] = comfy.conds.CONDRegular(torch.FloatTensor([guidance])) return out +class MiniMaxH3(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.model.MiniMaxH3Model) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + # run condition_proj + token refiner once per sampling instead of per step + cross_attn = self.diffusion_model.preprocess_text_embeds( + cross_attn.to(device=kwargs["device"], dtype=self.get_dtype_inference())) + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + + latent_shapes = kwargs.get("latent_shapes", None) + if latent_shapes is not None: + out['latent_shapes'] = comfy.conds.CONDConstant(latent_shapes) + + # Everything H3-specific rides in one dict so _apply_model's dtype cast + # (which would flatten fp32 cond latents and long tags to bf16) skips it. + payload = {} + tags = kwargs.get("minimax_token_tags", None) + if tags is not None: + payload["text_token_tags"] = tags + 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] + 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] + 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: + payload["audio_cond_noise_aug"] = kwargs["minimax_audio_cond_noise_aug"] + payload["seed"] = kwargs.get("seed", 0) + 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")) + out['minimax_payload'] = comfy.conds.CONDConstant(payload) + return out + + def scale_latent_inpaint(self, sigma, noise, latent_image, **kwargs): + return latent_image + 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) @@ -2213,10 +2281,7 @@ class Omnigen2(BaseModel): out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) ref_latents = kwargs.get("reference_latents", None) if ref_latents is not None: - latents = [] - for lat in ref_latents: - latents.append(self.process_latent_in(lat)) - out['ref_latents'] = comfy.conds.CONDList(latents) + out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents]) return out def extra_conds_shapes(self, **kwargs): @@ -2232,8 +2297,8 @@ class Boogu(Omnigen2): self.memory_usage_factor_conds = ("ref_latents",) class QwenImage(BaseModel): - def __init__(self, model_config, model_type=ModelType.FLUX, device=None): - super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel) + def __init__(self, model_config, model_type=ModelType.FLUX, device=None, unet_model=comfy.ldm.qwen_image.model.QwenImageTransformer2DModel): + super().__init__(model_config, model_type, device=device, unet_model=unet_model) self.memory_usage_factor_conds = ("ref_latents",) def extra_conds(self, **kwargs): @@ -2263,6 +2328,43 @@ class QwenImage(BaseModel): out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) return out +class MageFlow(QwenImage): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.mage_flow.model.MageFlowTransformer2DModel) + + def process_timestep(self, timestep, **kwargs): + # Mage runs in bf16 and rounds its timestep frequency table to the timestep dtype, keep that on fp32 devices. + return timestep.to(torch.bfloat16) + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 128, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 128]) + return out + +class JoyImage(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLOW, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel) + self.memory_usage_factor_conds = ("ref_latents",) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents]) + return out + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) + return out + class Ideogram4(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel) @@ -2278,6 +2380,35 @@ class Ideogram4(BaseModel): out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) return out +class Krea2(BaseModel): + def __init__(self, model_config, model_type=ModelType.FLUX, device=None): + super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT) + self.memory_usage_factor_conds = ("ref_latents",) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + cross_attn = kwargs.get("cross_attn", None) + if cross_attn is not None: + out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + latents = [] + for lat in ref_latents: + latents.append(self.process_latent_in(lat)) + out['ref_latents'] = comfy.conds.CONDList(latents) + + ref_latents_method = kwargs.get("reference_latents_method", None) + if ref_latents_method is not None: + out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method) + return out + + def extra_conds_shapes(self, **kwargs): + out = {} + ref_latents = kwargs.get("reference_latents", None) + if ref_latents is not None: + out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16]) + return out + class HunyuanImage21(BaseModel): def __init__(self, model_config, model_type=ModelType.FLOW, device=None): super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.hunyuan_video.model.HunyuanVideo) diff --git a/comfy/model_detection.py b/comfy/model_detection.py index b773f0393..103680fd1 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -359,6 +359,35 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): # PixArt diffusers return None + if '{}video_patch_proj.weight'.format(key_prefix) in state_dict_keys and '{}audio_patch_proj.weight'.format(key_prefix) in state_dict_keys: # MiniMax H3 + dit_config = {} + dit_config["image_model"] = "minimax_h3" + dit_config["num_layers"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') + dit_config["token_refiner_num_layers"] = count_blocks(state_dict_keys, '{}token_refiner.blocks.'.format(key_prefix) + '{}.') + dit_config["hidden_size"] = state_dict['{}video_patch_proj.weight'.format(key_prefix)].shape[0] + dit_config["latents_dim"] = state_dict['{}final_layer.video_out.weight'.format(key_prefix)].shape[0] // 4 # patch 1x2x2 + dit_config["audio_latents_dim"] = state_dict['{}final_layer.audio_out.weight'.format(key_prefix)].shape[0] + dit_config["attention_head_dim"] = state_dict['{}blocks.0.attn.q_norm.weight'.format(key_prefix)].shape[0] + qkv = state_dict['{}blocks.0.attn.qkv_proj.weight'.format(key_prefix)] + dit_config["num_attention_heads"] = qkv.shape[0] // (3 * dit_config["attention_head_dim"]) + dit_config["ffn_hidden_size"] = state_dict['{}blocks.0.mlp.fc1.weight'.format(key_prefix)].shape[0] // 2 + dit_config["text_dim"] = state_dict['{}condition_proj.weight'.format(key_prefix)].shape[1] + table_key = '{}adaln_t_table'.format(key_prefix) + if table_key in state_dict_keys: + # adaln shipped over a precomputed curve basis: the adaln linears span a small shared basis of the time-embedding curve (no time embedder) + table = state_dict[table_key].shape # [grid, k] + dit_config["adaln_curve_grid"] = table[0] + dit_config["time_embed_dim"] = table[1] + else: + te = state_dict['{}time_embedder.proj_in.weight'.format(key_prefix)] + dit_config["timestep_input_dim"] = te.shape[1] + dit_config["time_embed_hidden_size"] = te.shape[0] + dit_config["time_embed_dim"] = state_dict['{}time_embedder.proj_out.weight'.format(key_prefix)].shape[0] + dit_config["rope_inv_freq_len"] = state_dict['{}rope.inv_freq'.format(key_prefix)].shape[0] + if metadata is not None and "config" in metadata: + dit_config.update(json.loads(metadata["config"]).get("transformer", {})) + return dit_config + if '{}adaln_single.emb.timestep_embedder.linear_1.bias'.format(key_prefix) in state_dict_keys: #Lightricks ltxv dit_config = {} dit_config["image_model"] = "ltxav" if f'{key_prefix}audio_adaln_single.linear.weight' in state_dict_keys else "ltxv" @@ -470,15 +499,46 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): # PiD (Pixel Diffusion Decoder). Must check BEFORE plain PixelDiT_T2I. _lq_w_key = '{}lq_proj.latent_proj.0.weight'.format(key_prefix) if _lq_w_key in state_dict_keys: - in_ch = int(state_dict[_lq_w_key].shape[1]) + latent_proj_in_channels = int(state_dict[_lq_w_key].shape[1]) + hidden_dim = int(state_dict[_lq_w_key].shape[0]) _gate_prefix = '{}lq_proj.gate_modules.'.format(key_prefix) num_gates = len({k[len(_gate_prefix):].split('.')[0] for k in state_dict_keys if k.startswith(_gate_prefix)}) + pid_v1_5 = '{}lq_proj.pit_head.weight'.format(key_prefix) in state_dict_keys dit_config = {"image_model": "pid", - "lq_latent_channels": in_ch, - "latent_spatial_down_factor": 16 if in_ch >= 64 else 8} + "lq_hidden_dim": hidden_dim} if num_gates > 0: dit_config["lq_interval"] = (14 + num_gates - 1) // num_gates + if pid_v1_5: + pid_v1_5_variants = { + 16: { # Flux and QwenImage + "lq_latent_channels": 16, + "latent_spatial_down_factor": 8, + "lq_latent_unpatchify_factor": 1, + }, + 32: { # Flux2 after 2x latent unpatchify + "lq_latent_channels": 128, + "latent_spatial_down_factor": 16, + "lq_latent_unpatchify_factor": 2, + }, + } + variant = pid_v1_5_variants.get(latent_proj_in_channels) + if variant is None: + raise ValueError(f"Unsupported PiD v1.5 latent projection with {latent_proj_in_channels} input channels") + gate_weight = state_dict['{}lq_proj.gate_modules.0.content_proj.weight'.format(key_prefix)] + dit_config.update(variant) + dit_config.update({ + "lq_conv_padding_mode": "replicate", + "lq_gate_per_token": gate_weight.shape[0] == 1, + "pit_lq_inject": True, + "rope_ref_h": 2048, + "rope_ref_w": 2048, + }) + else: + dit_config.update({ + "lq_latent_channels": latent_proj_in_channels, + "latent_spatial_down_factor": 16 if latent_proj_in_channels >= 64 else 8, + }) return dit_config if '{}core.pixel_embedder.proj.weight'.format(key_prefix) in state_dict_keys: # PixelDiT T2I @@ -598,6 +658,44 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): return dit_config + seedvr2_7b_separate_key = "{}blocks.35.mlp.vid.proj_out.weight".format(key_prefix) + if seedvr2_7b_separate_key in state_dict_keys and state_dict[seedvr2_7b_separate_key].shape[0] == 3072: # seedvr2 7b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 3072 + dit_config["heads"] = 24 + dit_config["num_layers"] = 36 + # This checkpoint uses separate vid/txt MMModule keys in every block. + dit_config["mm_layers"] = 36 + dit_config["norm_eps"] = 1e-5 + dit_config["rope_type"] = "rope3d" + dit_config["rope_dim"] = 64 + dit_config["mlp_type"] = "normal" + return dit_config + if "{}blocks.35.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 7b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 3072 + dit_config["heads"] = 24 + dit_config["num_layers"] = 36 + # This checkpoint uses shared all.* MMModule keys after the initial blocks. + dit_config["mm_layers"] = 10 + dit_config["norm_eps"] = 1e-5 + dit_config["rope_type"] = "rope3d" + dit_config["rope_dim"] = 64 + dit_config["mlp_type"] = "swiglu" + return dit_config + if "{}blocks.31.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 3b + dit_config = {} + dit_config["image_model"] = "seedvr2" + dit_config["vid_dim"] = 2560 + dit_config["heads"] = 20 + dit_config["num_layers"] = 32 + dit_config["norm_eps"] = 1.0e-05 + dit_config["mlp_type"] = "swiglu" + dit_config["vid_out_norm"] = True + return dit_config + if '{}head.modulation'.format(key_prefix) in state_dict_keys: # Wan 2.1 dit_config = {} dit_config["image_model"] = "wan2.1" @@ -815,6 +913,13 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): "selected_layer_index": selected_layer_index, } + if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys and '{}proj_out.weight'.format(key_prefix) in state_dict_keys and state_dict['{}txt_norm.weight'.format(key_prefix)].shape[0] == 2560 and state_dict['{}proj_out.weight'.format(key_prefix)].shape[0] == 128: # Mage-Flow (Qwen Image txt_norm/proj_out are 3584/64) + dit_config = {} + dit_config["image_model"] = "mage_flow" + dit_config["in_channels"] = 128 + dit_config["num_layers"] = count_blocks(state_dict_keys, '{}transformer_blocks.'.format(key_prefix) + '{}.') + return dit_config + if '{}txt_norm.weight'.format(key_prefix) in state_dict_keys: # Qwen Image dit_config = {} dit_config["image_model"] = "qwen_image" @@ -834,6 +939,21 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["num_layers"] = count_blocks(state_dict_keys, '{}layers.'.format(key_prefix) + '{}.') return dit_config + if '{}txtfusion.projector.weight'.format(key_prefix) in state_dict_keys: # Krea 2 (K2) + dit_config = {} + dit_config["image_model"] = "krea2" + head_dim = 128 + first_w = state_dict['{}first.weight'.format(key_prefix)] # (features, channels*patch^2) + dit_config["features"] = first_w.shape[0] + dit_config["channels"] = first_w.shape[1] // (2 * 2) # patch=2 + dit_config["patch"] = 2 + dit_config["layers"] = count_blocks(state_dict_keys, '{}blocks.'.format(key_prefix) + '{}.') + dit_config["heads"] = state_dict['{}blocks.0.attn.wq.weight'.format(key_prefix)].shape[0] // head_dim + dit_config["kvheads"] = state_dict['{}blocks.0.attn.wk.weight'.format(key_prefix)].shape[0] // head_dim + dit_config["txtlayers"] = state_dict['{}txtfusion.projector.weight'.format(key_prefix)].shape[1] + dit_config["txtdim"] = state_dict['{}txtfusion.layerwise_blocks.0.prenorm.scale'.format(key_prefix)].shape[0] + return dit_config + if '{}visual_transformer_blocks.0.cross_attention.key_norm.weight'.format(key_prefix) in state_dict_keys: # Kandinsky 5 dit_config = {} model_dim = state_dict['{}visual_embeddings.in_layer.bias'.format(key_prefix)].shape[0] @@ -974,6 +1094,25 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): dit_config["image_model"] = "SAM31" return dit_config + if ( + '{}double_blocks.0.attn.img_attn_qkv.weight'.format(key_prefix) in state_dict_keys + and '{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix) in state_dict_keys + and '{}condition_embedder.time_embedder.linear_1.weight'.format(key_prefix) in state_dict_keys + and '{}img_in.weight'.format(key_prefix) in state_dict_keys + and len(state_dict['{}img_in.weight'.format(key_prefix)].shape) == 5 + ): + img_in = state_dict['{}img_in.weight'.format(key_prefix)] + head_dim = state_dict['{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix)].shape[0] + return { + "image_model": "joyimage", + "in_channels": img_in.shape[1], + "hidden_size": img_in.shape[0], + "patch_size": list(img_in.shape[2:]), + "num_layers": count_blocks(state_dict_keys, '{}double_blocks.'.format(key_prefix) + '{}.'), + "num_attention_heads": img_in.shape[0] // head_dim, + "text_dim": 4096, + } + if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys: return None @@ -1104,9 +1243,10 @@ def detect_unet_config(state_dict, key_prefix, metadata=None): return unet_config -def model_config_from_unet_config(unet_config, state_dict=None): + +def model_config_from_unet_config(unet_config, state_dict=None, unet_key_prefix=""): for model_config in comfy.supported_models.models: - if model_config.matches(unet_config, state_dict): + if model_config.matches(unet_config, state_dict, unet_key_prefix=unet_key_prefix): return model_config(unet_config) logging.error("no match {}".format(unet_config)) @@ -1116,7 +1256,7 @@ def model_config_from_unet(state_dict, unet_key_prefix, use_base_if_no_match=Fal unet_config = detect_unet_config(state_dict, unet_key_prefix, metadata=metadata) if unet_config is None: return None - model_config = model_config_from_unet_config(unet_config, state_dict) + model_config = model_config_from_unet_config(unet_config, state_dict, unet_key_prefix) if model_config is None and use_base_if_no_match: model_config = comfy.supported_models_base.BASE(unet_config) diff --git a/comfy/model_management.py b/comfy/model_management.py index f2569b0ea..b9c027c15 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -34,6 +34,7 @@ import comfy.utils import comfy.quant_ops import comfy_aimdo.host_buffer import comfy_aimdo.vram_buffer +from comfy.internal_logging import detail from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -481,7 +482,7 @@ except: SUPPORT_FP8_OPS = args.supports_fp8_compute -AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"] +AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"] AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN' try: @@ -624,6 +625,8 @@ PIN_PRESSURE_HYSTERESIS = 256 * 1024 * 1024 #Freeing registerables on pressure does imply a GPU sync, so go big on #the hysteresis so each expensive sync gives us back a good chunk. REGISTERABLE_PIN_HYSTERESIS = 2048 * 1024 * 1024 +WINDOWS_PIN_EVICTION_SWAP_PERCENT = 5.0 +WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE = 512 * 1024 ** 2 def module_size(module): module_mem = 0 @@ -638,19 +641,77 @@ def mark_mmap_dirty(storage): if mmap_refs is not None: DIRTY_MMAPS.add(mmap_refs[0]) -def free_pins(size, evict_active=False): +PIN_SUBSETS = [ "weights", "patches" ] +LOADED_PIN_SUBSETS = [ "weights-loaded", "patches-loaded" ] + +def models_for_pin_eviction(active, current_prompt=None): + for loaded_model in current_loaded_models: + model = loaded_model.model + if model is None or not model.is_dynamic(): + continue + pin_state = model.model.dynamic_pins[model.load_device] + if ((active is None or pin_state["active"] == active) and + (current_prompt is None or pin_state["current_prompt"] == current_prompt)): + yield model + +def free_model_pins(size, subsets, current_prompt, active, registrations=False): freed_total = 0 - for loaded_model in reversed(current_loaded_models): + for model in models_for_pin_eviction(active, current_prompt=current_prompt): if size <= 0: return freed_total - model = loaded_model.model - if model is not None and model.is_dynamic() and (evict_active or not model.model.dynamic_pins[model.load_device]["active"]): - freed = model.partially_unload_ram(size) - freed_total += freed - size -= freed + if registrations: + freed = model.unregister_inactive_pins(size, subsets=subsets) + else: + freed = model.partially_unload_ram(size, subsets=subsets) + freed_total += freed + size -= freed return freed_total -def ensure_pin_budget(size, evict_active=False): +def pin_eviction_tiers(loaded, evict_active): + tiers = [ + (PIN_SUBSETS, False, None), + (LOADED_PIN_SUBSETS, False, None), + (LOADED_PIN_SUBSETS, True, None), + ] + if not loaded: + tiers.append((PIN_SUBSETS, True, False)) + if evict_active: + tiers.append((PIN_SUBSETS, True, True)) + return tiers + +def registration_eviction_tiers(evict_active): + subsets = PIN_SUBSETS + LOADED_PIN_SUBSETS + tiers = [ + (subsets, False, False), + (subsets, True, False), + ] + if evict_active: + tiers.extend([ + (subsets, False, True), + (subsets, True, True), + ]) + return tiers + +def free_pins(size, evict_active=False, loaded=False): + freed = 0 + for subsets, current_prompt, active in pin_eviction_tiers(loaded, evict_active): + freed += free_model_pins(size - freed, subsets, current_prompt, active) + return freed + +def should_free_pins_for_ram_pressure(shortfall): + if shortfall <= 0: + return False + if not WINDOWS: + return True + if psutil.virtual_memory().available < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE: + return True + try: + return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT + except RuntimeError as err: + logging.warning("Could not read Windows swap usage; falling back to RAM-pressure pin eviction: %s", err) + return True + +def ensure_pin_budget(size, evict_active=False, loaded=False): if args.high_ram: return True if args.fast_disk: @@ -661,7 +722,7 @@ def ensure_pin_budget(size, evict_active=False): return True to_free = shortfall + PIN_PRESSURE_HYSTERESIS - return free_pins(to_free, evict_active=evict_active) >= shortfall + return free_pins(to_free, evict_active=evict_active, loaded=loaded) >= shortfall def free_registrations(shortfall, evict_active=True): if MAX_PINNED_MEMORY <= 0: @@ -670,19 +731,8 @@ def free_registrations(shortfall, evict_active=True): return True shortfall += REGISTERABLE_PIN_HYSTERESIS - for loaded_model in reversed(current_loaded_models): - model = loaded_model.model - if model is not None and model.is_dynamic() and not model.model.dynamic_pins[model.load_device]["active"]: - shortfall -= model.unregister_inactive_pins(shortfall) - if shortfall <= 0: - return True - if evict_active: - for loaded_model in current_loaded_models: - model = loaded_model.model - if model is not None and model.is_dynamic() and model.model.dynamic_pins[model.load_device]["active"]: - shortfall -= model.unregister_inactive_pins(shortfall) - if shortfall <= 0: - return True + for subsets, current_prompt, active in registration_eviction_tiers(evict_active): + shortfall -= free_model_pins(shortfall, subsets, current_prompt, active, registrations=True) return shortfall <= REGISTERABLE_PIN_HYSTERESIS def ensure_pin_registerable(size, evict_active=True): @@ -812,6 +862,8 @@ def minimum_inference_memory(): def free_memory(memory_required, device, keep_loaded=[], for_dynamic=False, pins_required=0, ram_required=0): cleanup_models_gc() + if not for_dynamic: + detail("Non dynamic memory free called! memory_required=%s pins_required=%s ram_required=%s", memory_required, pins_required, ram_required) unloaded_model = [] can_unload = [] unloaded_models = [] @@ -957,6 +1009,9 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu ) loaded_model.model_load(lowvram_model_memory, force_patch_weights=force_patch_weights) logging.info(f"Loading model {model_name} complete") + vram_used = 0 if is_device_cpu(torch_dev) else loaded_model.model_loaded_memory() + ram_used = model.loaded_ram_size() if model.is_dynamic() else loaded_model.model_memory() - vram_used + detail("Model loaded: patcher=%s model=%s ram_mb=%.1f vram_mb=%.1f", model.__class__.__name__, model.model.__class__.__name__, ram_used / (1024 ** 2), vram_used / (1024 ** 2)) current_loaded_models.insert(0, loaded_model) return @@ -1383,15 +1438,17 @@ def reset_cast_buffers(): pin_state = model.model.dynamic_pins[model.load_device] if pin_state["active"]: - *_, buckets = pin_state["weights"] - for size, bucket in list(buckets.items()): - bucket[:] = [ entry for entry in bucket if entry[-1] is not None ] - if not bucket: - del buckets[size] + for subset in ("weights", "weights-loaded"): + *_, buckets = pin_state[subset] + for size, bucket in list(buckets.items()): + bucket[:] = [ entry for entry in bucket if entry[-1] is not None ] + if not bucket: + del buckets[size] pin_state["active"] = False - model.partially_unload_ram(1e30, subsets=[ "patches" ]) - model.model.dynamic_pins[model.load_device]["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {}) + model.partially_unload_ram(1e30, subsets=[ "patches", "patches-loaded" ]) + for subset in ("patches", "patches-loaded"): + pin_state[subset] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, pinned_hostbuf_size(model.model_size())), [], [-1], [0], [0], {}) STREAM_CAST_BUFFERS.clear() STREAM_AIMDO_CAST_BUFFERS.clear() diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index 2814563db..e7d2d7727 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -22,6 +22,7 @@ import collections import inspect import logging import math +import time import uuid from typing import Callable, Optional @@ -37,12 +38,58 @@ import comfy.patcher_extension import comfy.utils import comfy_aimdo.host_buffer from comfy.comfy_types import UnetWrapperFunction +from comfy.internal_logging import detail from comfy.quant_ops import QuantizedTensor from comfy.patcher_extension import CallbacksMP, PatcherInjection, WrappersMP import comfy_aimdo.model_vbar _WSL_MODEL_LOAD_SYNC_SKIP_LOGGED = False +def is_model_patcher_output(output): + return isinstance(output, ModelPatcher) or isinstance(getattr(output, "patcher", None), ModelPatcher) + +class PromptModelTracker: + def __init__(self): + self.models = {} + + def start(self): + self.end() + + def add(self, outputs): + if isinstance(outputs, collections.abc.Mapping): + outputs = outputs.values() + elif not isinstance(outputs, (list, tuple)): + outputs = (outputs,) + + for output in outputs: + if isinstance(output, (collections.abc.Mapping, list, tuple)): + self.add(output) + continue + + models = [] + if isinstance(output, ModelPatcher): + models.append(output) + models.extend(output.model_patches_models()) + models.extend(output.get_nested_additional_models()) + else: + patcher = getattr(output, "patcher", None) + if isinstance(patcher, ModelPatcher): + models.append(patcher) + get_models = getattr(output, "get_models", None) + if callable(get_models): + models.extend(get_models()) + + for model in models: + if not isinstance(model, ModelPatcher) or not model.is_dynamic(): + continue + key = (id(model.model), model.load_device) + self.models[key] = model + model.set_in_use_by_current_prompt(True) + + def end(self): + for model in self.models.values(): + model.set_in_use_by_current_prompt(False) + self.models.clear() def set_model_options_patch_replace(model_options, patch, name, block_name, number, transformer_index=None): to = model_options["transformer_options"].copy() @@ -512,12 +559,9 @@ class ModelPatcher: new_multigpu_models = [] for mm in multigpu_models: # clone main model, but bring over relevant props from existing multigpu clone - n = self.clone() + n = self.clone(model_override=mm.get_clone_model_override()) n.load_device = mm.load_device - n.backup = mm.backup - n.object_patches_backup = mm.object_patches_backup n.hook_backup = mm.hook_backup - n.model = mm.model n.is_multigpu_base_clone = mm.is_multigpu_base_clone n.remove_additional_models("multigpu") orig_additional_models: dict[str, list[ModelPatcher]] = comfy.patcher_extension.copy_nested_dicts(n.additional_models) @@ -1717,6 +1761,9 @@ class ModelPatcherDynamic(ModelPatcher): self.register_load_device(self.load_device) self.non_dynamic_delegate_model = None assert load_device is not None + if not hasattr(self.model, "dynamic_patchers"): + self.model.dynamic_patchers = set() + self.model.dynamic_patchers.add(id(self)) def register_load_device(self, device): """Ensure dynamic_pins has an entry for *device*. @@ -1731,14 +1778,20 @@ class ModelPatcherDynamic(ModelPatcher): self.model.dynamic_pins[device] = { "weights": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), "patches": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), + "weights-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), + "patches-loaded": (comfy_aimdo.host_buffer.HostBuffer(0, 0, 0), [], [-1], [0], [0], {}), "hostbufs_initialized": False, "failed": False, "active": False, + "current_prompt": False, } def is_dynamic(self): return True + def set_in_use_by_current_prompt(self, in_use): + self.model.dynamic_pins[self.load_device]["current_prompt"] = in_use + def _vbar_get(self, create=False): if self.load_device == torch.device("cpu"): return None @@ -1766,6 +1819,18 @@ class ModelPatcherDynamic(ModelPatcher): def unpin_all_weights(self): self.partially_unload_ram(1e32) + def __del__(self): + model = getattr(self, "model", None) + dynamic_patchers = getattr(model, "dynamic_patchers", None) + if dynamic_patchers is None or id(self) not in dynamic_patchers: + return + dynamic_patchers.discard(id(self)) + try: + if not dynamic_patchers: + self.unpin_all_weights() + finally: + self.detach(unpatch_all=False) + def memory_required(self, input_shape): #Pad this significantly. We are trying to get away from precise estimates. This #estimate is only used when using the ModelPatcherDynamic after ModelPatcher. If you @@ -1809,6 +1874,8 @@ class ModelPatcherDynamic(ModelPatcher): hostbuf_size = comfy.model_management.pinned_hostbuf_size(self.model_size()) pin_state["weights"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) pin_state["patches"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) + pin_state["weights-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 64 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) + pin_state["patches-loaded"] = (comfy_aimdo.host_buffer.HostBuffer(0, 8 * 1024 * 1024, hostbuf_size), [], [-1], [0], [0], {}) pin_state["hostbufs_initialized"] = True pin_state["failed"] = False pin_state["active"] = True @@ -1942,20 +2009,37 @@ class ModelPatcherDynamic(ModelPatcher): assert self.load_device != torch.device("cpu") vbar = self._vbar_get() - freed = 0 if vbar is None else vbar.free_memory(memory_to_free) + vbar_freed = 0 if vbar is None else vbar.free_memory(memory_to_free) + freed = vbar_freed + backup_freed = 0 if freed < memory_to_free: - freed += self.restore_loaded_backups() + backup_freed = self.restore_loaded_backups() + freed += backup_freed + + method = "vbar+backups" if vbar_freed and backup_freed else "vbar" if vbar_freed else "backups" if backup_freed else "none" + free_methods = getattr(self, "_free_methods", {}) + free_methods[method] = free_methods.get(method, 0) + 1 + self._free_methods = free_methods + now = time.monotonic() + if now - getattr(self, "_last_free_log_time", 0) >= 5: + requested = "all" if memory_to_free >= 1e30 else f"{memory_to_free / (1024 ** 2):.1f}MB" + prevailing_method = max(free_methods, key=free_methods.get) + detail("AIMDO free: model=%s device=%s prevailing_method=%s methods=%s requested=%s vbar_mb=%.1f backups_mb=%.1f", self.model.__class__.__name__, self.load_device, prevailing_method, free_methods, requested, vbar_freed / (1024 ** 2), backup_freed / (1024 ** 2)) + self._free_methods = {} + self._last_free_log_time = now return freed def loaded_ram_size(self): - return (self.model.dynamic_pins[self.load_device]["weights"][0].size) + pin_state = self.model.dynamic_pins[self.load_device] + return pin_state["weights"][0].size + pin_state["weights-loaded"][0].size def pinned_memory_size(self): - return (self.model.dynamic_pins[self.load_device]["weights"][3][0]) + pin_state = self.model.dynamic_pins[self.load_device] + return pin_state["weights"][3][0] + pin_state["weights-loaded"][3][0] - def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights", "patches" ]): + def unregister_inactive_pins(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]): freed = 0 pin_state = self.model.dynamic_pins[self.load_device] for subset in subsets: @@ -1963,15 +2047,17 @@ class ModelPatcherDynamic(ModelPatcher): split = stack_split[0] while split >= 0: module, offset = stack[split] + module_pin = module._pins[subset] split -= 1 stack_split[0] = split - if not module._pin_registered: + if not module_pin["registered"]: continue - size = module._pin.numel() * module._pin.element_size() - if torch.cuda.cudart().cudaHostUnregister(module._pin.data_ptr()) != 0: + pin = module_pin["pin"] + size = pin.numel() * pin.element_size() + if torch.cuda.cudart().cudaHostUnregister(pin.data_ptr()) != 0: comfy.model_management.discard_cuda_async_error() continue - module._pin_registered = False + module_pin["registered"] = False comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size) pinned_size[0] = max(0, pinned_size[0] - size) freed += size @@ -1980,20 +2066,23 @@ class ModelPatcherDynamic(ModelPatcher): return freed return freed - def partially_unload_ram(self, ram_to_unload, subsets=[ "weights", "patches" ]): + def partially_unload_ram(self, ram_to_unload, subsets=[ "weights-loaded", "patches-loaded", "weights", "patches" ]): freed = 0 pin_state = self.model.dynamic_pins[self.load_device] for subset in subsets: hostbuf, stack, stack_split, pinned_size, *_ = pin_state[subset] while len(stack) > 0: module, offset = stack.pop() - size = module._pin.numel() * module._pin.element_size() - module._pin_balancer_entry[-1] = None - del module._pin_balancer_entry - del module._pin - hostbuf.truncate(offset, do_unregister=module._pin_registered) + module_pin = module._pins[subset] + pin = module_pin["pin"] + size = pin.numel() * pin.element_size() + module_pin["balancer_entry"][-1] = None + del module_pin["balancer_entry"] + del module_pin["pin"] + registered = module_pin["registered"] + hostbuf.truncate(offset, do_unregister=registered) stack_split[0] = min(stack_split[0], len(stack) - 1) - if module._pin_registered: + if registered: comfy.model_management.TOTAL_PINNED_MEMORY = max(0, comfy.model_management.TOTAL_PINNED_MEMORY - size) pinned_size[0] = max(0, pinned_size[0] - size) freed += size diff --git a/comfy/nested_tensor.py b/comfy/nested_tensor.py index b700816fa..08c7133f8 100644 --- a/comfy/nested_tensor.py +++ b/comfy/nested_tensor.py @@ -51,6 +51,9 @@ class NestedTensor: def float(self): return self.to(dtype=torch.float) + def cpu(self): + return self.to(device="cpu") + def chunk(self, *args, **kwargs): return self.apply_operation(None, lambda x, y: x.chunk(*args, **kwargs)) diff --git a/comfy/ops.py b/comfy/ops.py index 3f088a962..077a42351 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -19,6 +19,7 @@ import torch import logging import contextlib +import inspect import comfy.model_management from comfy.cli_args import args, PerformanceFeature import comfy.float @@ -36,27 +37,59 @@ def run_every_op(): comfy.model_management.throw_exception_if_processing_interrupted() +def gqa_repeat_factor(query_heads, key_heads, value_heads): + if key_heads != value_heads: + raise ValueError(f"Key/value head count mismatch for GQA: {key_heads} != {value_heads}") + if query_heads == key_heads: + return 1 + if query_heads % key_heads != 0: + raise ValueError(f"Query heads must be divisible by key/value heads for GQA: {query_heads} vs {key_heads}") + return query_heads // key_heads + +def repeat_kv_for_gqa(k, v, query_heads, head_dim): + n_rep = gqa_repeat_factor(query_heads, k.shape[head_dim], v.shape[head_dim]) + if n_rep > 1: + k = k.repeat_interleave(n_rep, dim=head_dim) + v = v.repeat_interleave(n_rep, dim=head_dim) + return k, v + def scaled_dot_product_attention(q, k, v, *args, **kwargs): + attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask") + if kwargs.get("enable_gqa", False) and attn_mask is not None: + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) try: - if torch.cuda.is_available() and comfy.model_management.WINDOWS: + if torch.cuda.is_available(): from torch.nn.attention import SDPBackend, sdpa_kernel - import inspect if "set_priority" in inspect.signature(sdpa_kernel).parameters: SDPA_BACKEND_PRIORITY = [ SDPBackend.FLASH_ATTENTION, + SDPBackend.CUDNN_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH, ] - SDPA_BACKEND_PRIORITY.insert(0, SDPBackend.CUDNN_ATTENTION) - def scaled_dot_product_attention(q, k, v, *args, **kwargs): - if q.nelement() < 1024 * 128: # arbitrary number, for small inputs cudnn attention seems slower - return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) + attn_mask = args[0] if len(args) > 0 else kwargs.get("attn_mask") + if kwargs.get("enable_gqa", False) and attn_mask is not None and not comfy.model_management.is_nvidia(): + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False with sdpa_kernel(SDPA_BACKEND_PRIORITY, set_priority=True): + if kwargs.get("enable_gqa", False) and attn_mask is not None and q.shape[-3] != k.shape[-3]: + dropout_p = args[1] if len(args) > 1 else kwargs.get("dropout_p", 0.0) + is_causal = args[2] if len(args) > 2 else kwargs.get("is_causal", False) + params = torch.backends.cuda.SDPAParams(q, k, v, attn_mask, dropout_p, is_causal, True) + supports_native_gqa = ( + torch.backends.cuda.can_use_flash_attention(params) + or torch.backends.cuda.can_use_cudnn_attention(params) + or torch.backends.cuda.can_use_efficient_attention(params) + ) + if not supports_native_gqa: + k, v = repeat_kv_for_gqa(k, v, q.shape[-3], -3) + kwargs["enable_gqa"] = False return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) else: logging.warning("Torch version too old to set sdpa backend priority.") @@ -144,8 +177,13 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin needs_cast = False xfer_source = [ s.weight, s.bias ] - - pin = comfy.pinned_memory.get_pin(s) + subset = "weights" + pin = comfy.pinned_memory.get_pin(s, subset=subset) + if pin is None and not args.fast_disk: + loaded_pin = comfy.pinned_memory.get_pin(s, subset="weights-loaded") + if loaded_pin is not None or signature is not None: + subset = "weights-loaded" + pin = loaded_pin if pin is not None: xfer_source = [ pin ] @@ -174,18 +212,20 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin elif xfer_dest2 is not None: xfer_source.prepare(xfer_dest2, stream, copy=True, commit=False) return + else: + return comfy.model_management.cast_to_gathered(xfer_source, xfer_dest, non_blocking=non_blocking, stream=stream, r2=xfer_dest2) def handle_pin(m, pin, source, dest, subset="weights", size=None): if pin is not None: cast_maybe_lowvram_patch([pin], dest, offload_stream) return - if signature is None or args.high_ram: + if signature is None or not args.fast_disk or args.high_ram: comfy.pinned_memory.pin_memory(m, subset=subset, size=size) pin = comfy.pinned_memory.get_pin(m, subset=subset) cast_maybe_lowvram_patch(source, pin, offload_stream, xfer_dest2=dest) - handle_pin(s, pin, xfer_source, xfer_dest, size=dest_size) + handle_pin(s, pin, xfer_source, xfer_dest, subset=subset, size=dest_size) for param_key in ("weight", "bias"): lowvram_source = getattr(s, param_key + "_lowvram_function", None) @@ -195,8 +235,16 @@ def cast_modules_with_vbar(comfy_modules, dtype, device, bias_dtype, non_blockin lowvram_dest = get_cast_buffer(lowvram_size) lowvram_source.prepare(lowvram_dest, None, copy=False, commit=True) - pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches") - handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset="patches", size=lowvram_size) + subset = "patches" + pin = comfy.pinned_memory.get_pin(lowvram_source, subset=subset) + if pin is None: + loaded_pin = comfy.pinned_memory.get_pin(lowvram_source, subset="patches-loaded") + if loaded_pin is not None: + subset = "patches-loaded" + pin = loaded_pin + elif signature is not None and not args.fast_disk: + subset = "patches-loaded" + handle_pin(lowvram_source, pin, lowvram_source, lowvram_dest, subset=subset, size=lowvram_size) prefetch["xfer_dest"] = xfer_dest @@ -256,7 +304,7 @@ def resolve_cast_module_with_vbar(s, dtype, device, bias_dtype, compute_dtype, w if (want_requant and len(fns) == 0 or update_weight): seed = comfy.utils.string_to_seed(s.seed_key) if isinstance(orig, QuantizedTensor): - y = QuantizedTensor.from_float(x, s.layout_type, scale="recalculate", stochastic_rounding=seed) + y = orig.requantize_from_float(x, scale="recalculate", stochastic_rounding=seed) else: y = comfy.float.stochastic_rounding(x, orig.dtype, seed=seed) if want_requant and len(fns) == 0: @@ -446,8 +494,7 @@ class disable_weight_init: def __init__(self, in_features, out_features, bias=True, device=None, dtype=None): # don't trust subclasses that BYO state dict loader to call us. - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Linear._load_from_state_dict): super().__init__(in_features, out_features, bias, device, dtype) return @@ -469,8 +516,7 @@ class disable_weight_init: def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs): - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Linear._load_from_state_dict): return super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) @@ -698,8 +744,7 @@ class disable_weight_init: norm_type=2.0, scale_grad_by_freq=False, sparse=False, _weight=None, _freeze=False, device=None, dtype=None): # don't trust subclasses that BYO state dict loader to call us. - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Embedding._load_from_state_dict): super().__init__(num_embeddings, embedding_dim, padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse, _weight, @@ -726,8 +771,7 @@ class disable_weight_init: def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs): - if (not comfy.model_management.WINDOWS - or not comfy.memory_management.aimdo_enabled + if (not comfy.memory_management.aimdo_enabled or type(self)._load_from_state_dict is not disable_weight_init.Embedding._load_from_state_dict): return super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) @@ -899,13 +943,61 @@ if CUBLAS_IS_AVAILABLE: # ============================================================================== # Mixed Precision Operations # ============================================================================== +from . import quant_ops from .quant_ops import ( QuantizedTensor, QUANT_ALGOS, TensorCoreFP8Layout, + TensorWiseINT8Layout, get_layout_class, ) +def _swiglu_eager(x): + gate, up = x.chunk(2, dim=-1) + return torch.nn.functional.silu(gate).mul_(up) + + +INPUT_ACT_EAGER = { + "gelu_tanh": lambda x: torch.nn.functional.gelu(x, approximate="tanh"), + "swiglu": _swiglu_eager, +} + + +def linear_input_act(linear, x, input_act): + """``linear(act(x))``, with ``act`` folded into an INT8 activation quantizer. + + An INT8 linear quantizes its input anyway, so an elementwise activation can + ride along inside that kernel instead of writing a full-size intermediate to + HBM and reading it straight back. Worth it for an MLP's down-projection, + where the intermediate is several times the hidden size. + + """ + weight = linear.weight + if (comfy.model_management.in_training + or not isinstance(weight, QuantizedTensor) + or weight._layout_cls != "TensorWiseINT8Layout" + or getattr(weight._params, "transposed", False)): + return linear(INPUT_ACT_EAGER[input_act](x)) + + # want_requant keeps a vbar-streamed layer on the INT8 path when a LoRA is + # patched in on the fly; without it the cast hands back a dequantized weight. + weight, bias, offload_stream = cast_bias_weight( + linear, x, offloadable=True, compute_dtype=x.dtype, want_requant=True) + try: + if not isinstance(weight, QuantizedTensor): + # A LoRA weight_function, or activations whose dtype differs from the + # weight's, make the cast hand back a dequantized tensor. + return torch.nn.functional.linear(INPUT_ACT_EAGER[input_act](x), weight, bias) + qdata, scale = TensorWiseINT8Layout.get_plain_tensors(weight) + return quant_ops.ck.int8_linear( + x, qdata, scale, bias, x.dtype, + convrot=getattr(weight._params, "convrot", False), + convrot_groupsize=getattr(weight._params, "convrot_groupsize", 256), + input_act=input_act, + ) + finally: + uncast_bias_weight(linear, weight, bias, offload_stream) + class QuantLinearFunc(torch.autograd.Function): """Custom autograd function for quantized linear: quantized forward, optionally FP8 backward. @@ -1089,6 +1181,34 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat if ts is None or bs is None: raise ValueError(f"Missing NVFP4 scales for layer {layer_name}") scales = {"scale": ts, "block_scale": bs} + elif module.quant_format == "int8_tensorwise": + scale = pop_scale("weight_scale") + if scale is None: + raise ValueError(f"Missing INT8 weight scale for layer {layer_name}") + scales = {"scale": scale} + params_conf = layer_conf.get("params", {}) + if not isinstance(params_conf, dict): + params_conf = {} + if layer_conf.get("convrot", params_conf.get("convrot", False)): + scales["convrot"] = True + scales["convrot_groupsize"] = int( + layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)) + ) + elif module.quant_format == "convrot_w4a4": + scale = pop_scale("weight_scale") + if scale is None: + raise ValueError(f"Missing ConvRot W4A4 weight scale for layer {layer_name}") + params_conf = layer_conf.get("params", {}) + if not isinstance(params_conf, dict): + params_conf = {} + scales = { + "scale": scale, + "convrot_groupsize": int( + layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)) + ), + "quant_group_size": 64, + "linear_dtype": layer_conf.get("linear_dtype", params_conf.get("linear_dtype", "int4")), + } else: raise ValueError(f"Unsupported quantization format: {module.quant_format}") @@ -1131,6 +1251,15 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr quant_conf = {"format": module.quant_format} if getattr(module, '_full_precision_mm_config', False): quant_conf["full_precision_matrix_mult"] = True + params = getattr(module.weight, "_params", None) + if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False): + quant_conf["convrot"] = True + quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256) + elif module.quant_format == "convrot_w4a4": + quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256) + linear_dtype = getattr(params, "linear_dtype", "int4") + if linear_dtype != "int4": + quant_conf["linear_dtype"] = linear_dtype if extra_quant_conf: quant_conf.update(extra_quant_conf) sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8) @@ -1178,13 +1307,38 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def state_dict(self, *args, destination=None, prefix="", **kwargs): sd = destination if destination is not None else {} - return _quantized_weight_state_dict(self, sd, prefix, extra_quant_params=("input_scale",)) + return _quantized_weight_state_dict(self, sd, prefix, extra_quant_params=("input_scale", "pre_quant_scale")) def _forward(self, input, weight, bias): return torch.nn.functional.linear(input, weight, bias) - def forward_comfy_cast_weights(self, input, compute_dtype=None, want_requant=False): - weight, bias, offload_stream = cast_bias_weight(self, input, offloadable=True, compute_dtype=compute_dtype, want_requant=want_requant) + def forward_comfy_cast_weights( + self, + input, + compute_dtype=None, + 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( + 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 @@ -1192,8 +1346,13 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def forward(self, input, *args, **kwargs): run_every_op() + # ModelOpt AWQ-style smoothing + pre_quant_scale = getattr(self, 'pre_quant_scale', None) + if pre_quant_scale is not None: + input = input * comfy.model_management.cast_to_device(pre_quant_scale, input.device, input.dtype) + input_shape = input.shape - reshaped_3d = False + reshaped_nd = False #If cast needs to apply lora, it should be done in the compute dtype compute_dtype = input.dtype @@ -1203,9 +1362,10 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec not getattr(self, 'comfy_force_cast_weights', False) and len(self.weight_function) == 0 and len(self.bias_function) == 0 ) + quantize_input = QUANT_ALGOS.get(getattr(self, 'quant_format', None), {}).get("quantize_input", True) # Training path: quantized forward with compute_dtype backward via autograd function - if (input.requires_grad and _use_quantized): + if (input.requires_grad and _use_quantized and quantize_input): weight, bias, offload_stream = cast_bias_weight( self, @@ -1227,25 +1387,31 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec return output # Inference path (unchanged) - if _use_quantized: + if _use_quantized and quantize_input: - # Reshape 3D tensors to 2D for quantization (needed for NVFP4 and others) - input_reshaped = input.reshape(-1, input_shape[2]) if input.ndim == 3 else input + # Reshape >=3D tensors to 2D for quantization (needed for NVFP4 and others) + input_reshaped = input.reshape(-1, input_shape[-1]) if input.ndim >= 3 else input # Fall back to non-quantized for non-2D tensors if input_reshaped.ndim == 2: - reshaped_3d = input.ndim == 3 + reshaped_nd = input.ndim >= 3 # dtype is now implicit in the layout class scale = getattr(self, 'input_scale', None) if scale is not None: scale = comfy.model_management.cast_to_device(scale, input.device, None) input = QuantizedTensor.from_float(input_reshaped, self.layout_type, scale=scale) - output = self.forward_comfy_cast_weights(input, compute_dtype, want_requant=isinstance(input, QuantizedTensor)) + weight_only_quant = _use_quantized and not quantize_input and isinstance(self.weight, QuantizedTensor) + output = self.forward_comfy_cast_weights( + input, + compute_dtype, + want_requant=isinstance(input, QuantizedTensor), + weight_only_quant=weight_only_quant, + ) - # Reshape output back to 3D if input was 3D - if reshaped_3d: - output = output.reshape((input_shape[0], input_shape[1], self.weight.shape[0])) + # Reshape output back to original rank if input was >2D + if reshaped_nd: + output = output.reshape((*input_shape[:-1], self.weight.shape[0])) return output @@ -1257,8 +1423,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def set_weight(self, weight, inplace_update=False, seed=None, return_weight=False, **kwargs): if getattr(self, 'layout_type', None) is not None: - # dtype is now implicit in the layout class - weight = QuantizedTensor.from_float(weight, self.layout_type, scale="recalculate", stochastic_rounding=seed, inplace_ops=True).to(self.weight.dtype) + weight = self.weight.requantize_from_float(weight, scale="recalculate", stochastic_rounding=seed, inplace_ops=True).to(self.weight.dtype) else: weight = weight.to(self.weight.dtype) if return_weight: @@ -1380,6 +1545,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec } if hasattr(params, "block_scale"): # NVFP4 kwargs["block_scale"] = params.block_scale[i] + if hasattr(params, "quant_group_size"): + kwargs["quant_group_size"] = params.quant_group_size + if hasattr(params, "convrot_groupsize"): + kwargs["convrot_groupsize"] = params.convrot_groupsize + if hasattr(params, "linear_dtype"): + kwargs["linear_dtype"] = params.linear_dtype return QuantizedTensor(weight._qdata[i], weight._layout_cls, type(params)(**kwargs)) def state_dict(self, *args, destination=None, prefix="", **kwargs): @@ -1393,12 +1564,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec if layer_conf is not None: layer_conf = json.loads(layer_conf.numpy().tobytes()) - # Only fp8 makes sense for embeddings (per-row dequant via index select). + # Only fp8 and int8_tensorwise support per-row dequant via index select. # Block-scaled formats (NVFP4, MXFP8) can't do per-row lookup efficiently. quant_format = layer_conf.get("format") if layer_conf is not None else None manually_loaded_keys = [] - if quant_format in ("float8_e4m3fn", "float8_e5m2") and weight_key in state_dict: + if quant_format in ("float8_e4m3fn", "float8_e5m2", "int8_tensorwise") and weight_key in state_dict: self.quant_format = quant_format qconfig = QUANT_ALGOS[quant_format] self.layout_type = qconfig["comfy_tensor_layout"] @@ -1412,10 +1583,16 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec scale = scale.float() manually_loaded_keys.append(scale_key) + extra = {} + if quant_format == "int8_tensorwise" and layer_conf.get("convrot", False): + # rotated embedding table: record it so the forward un-rotates after lookup + extra["convrot"] = True + extra["convrot_groupsize"] = int(layer_conf.get("convrot_groupsize", 256)) params = layout_cls.Params( scale=scale if scale is not None else torch.ones((), dtype=torch.float32), orig_dtype=MixedPrecisionOps._compute_dtype, orig_shape=(self.num_embeddings, self.embedding_dim), + **extra, ) self.weight = torch.nn.Parameter( QuantizedTensor(weight.to(dtype=qconfig["storage_t"]), qconfig["comfy_tensor_layout"], params), @@ -1437,15 +1614,23 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec def forward_comfy_cast_weights(self, input, out_dtype=None): weight = self.weight - # Optimized path: lookup in fp8, dequantize only the selected rows. + # 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): - scale = qdata._params.scale + 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) + x = torch.nn.functional.embedding( input, qdata, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse) diff --git a/comfy/pinned_memory.py b/comfy/pinned_memory.py index cb77c517a..e9a9a70e2 100644 --- a/comfy/pinned_memory.py +++ b/comfy/pinned_memory.py @@ -9,14 +9,14 @@ import torch from comfy.cli_args import args -def _add_to_bucket(module, buckets, size, priority): +def _add_to_bucket(module, module_pin, buckets, size, priority): bucket = buckets.setdefault(size, []) entry = [-priority, 0, module] entry[1] = id(entry) bisect.insort(bucket, entry) - module._pin_balancer_entry = entry + module_pin["balancer_entry"] = entry -def _steal_pin(module, stack, buckets, size, priority): +def _steal_pin(module, stack, buckets, size, priority, subset): bucket = buckets.get(size) if bucket is None: return False @@ -31,22 +31,27 @@ def _steal_pin(module, stack, buckets, size, priority): return False *_, victim = bucket.pop() - module._pin = victim._pin - module._pin_registered = victim._pin_registered - module._pin_stack_index = victim._pin_stack_index - stack[module._pin_stack_index] = (module, stack[module._pin_stack_index][1]) + module_pin = module._pins[subset] + victim_pin = victim._pins[subset] + module_pin["pin"] = victim_pin["pin"] + module_pin["registered"] = victim_pin["registered"] + module_pin["stack_index"] = victim_pin["stack_index"] + stack_index = module_pin["stack_index"] + stack[stack_index] = (module, stack[stack_index][1]) - victim._pin_registered = False - del victim._pin - del victim._pin_stack_index - del victim._pin_balancer_entry + victim_pin["registered"] = False + del victim_pin["pin"] + del victim_pin["stack_index"] + del victim_pin["balancer_entry"] - _add_to_bucket(module, buckets, size, priority) + _add_to_bucket(module, module_pin, buckets, size, priority) return True def get_pin(module, subset="weights"): - pin = getattr(module, "_pin", None) - if pin is None or module._pin_registered or args.disable_pinned_memory: + pins = module.__dict__.get("_pins") + module_pin = None if pins is None else pins.get(subset) + pin = None if module_pin is None else module_pin.get("pin") + if pin is None or module_pin["registered"] or args.disable_pinned_memory: return pin _, _, stack_split, pinned_size, *_ = module._pin_state[subset] @@ -57,8 +62,8 @@ def get_pin(module, subset="weights"): comfy.model_management.discard_cuda_async_error() return pin - module._pin_registered = True - stack_split[0] = max(stack_split[0], module._pin_stack_index) + module_pin["registered"] = True + stack_split[0] = max(stack_split[0], module_pin["stack_index"]) comfy.model_management.TOTAL_PINNED_MEMORY += size pinned_size[0] += size return pin @@ -72,23 +77,26 @@ def pin_memory(module, subset="weights", size=None): if pin is not None: return + pins = module.__dict__.setdefault("_pins", {}) + module_pin = pins.setdefault(subset, {}) hostbuf, stack, stack_split, pinned_size, counter, buckets = pin_state[subset] if size is None: size = comfy.memory_management.vram_aligned_size([ module.weight, module.bias ]) - offset = hostbuf.size registerable_size = size - priority = getattr(module, "_pin_balancer_priority", None) + loaded = subset.endswith("-loaded") + priority = module_pin.get("balancer_priority") if priority is None: priority = comfy.utils.bit_reverse_range(counter[0], 16) counter[0] += 1 - module._pin_balancer_priority = priority + module_pin["balancer_priority"] = priority comfy.memory_management.extra_ram_release(comfy.memory_management.RAM_CACHE_HEADROOM) - if (not comfy.model_management.ensure_pin_budget(size) or + if (not comfy.model_management.ensure_pin_budget(size, loaded=loaded) or not comfy.model_management.ensure_pin_registerable(registerable_size)): - return _steal_pin(module, stack, buckets, size, priority) + return _steal_pin(module, stack, buckets, size, priority, subset) + offset = hostbuf.size extended = False try: hostbuf.extend(size=size, register=False) @@ -102,18 +110,18 @@ def pin_memory(module, subset="weights", size=None): comfy.model_management.discard_cuda_async_error() del pin hostbuf.truncate(offset, do_unregister=False) - return _steal_pin(module, stack, buckets, size, priority) + return _steal_pin(module, stack, buckets, size, priority, subset) except RuntimeError: if extended: hostbuf.truncate(offset, do_unregister=False) - return _steal_pin(module, stack, buckets, size, priority) + return _steal_pin(module, stack, buckets, size, priority, subset) - module._pin = pin + module_pin["pin"] = pin stack.append((module, offset)) - module._pin_registered = True - module._pin_stack_index = len(stack) - 1 - stack_split[0] = max(stack_split[0], module._pin_stack_index) + module_pin["registered"] = True + module_pin["stack_index"] = len(stack) - 1 + stack_split[0] = max(stack_split[0], module_pin["stack_index"]) comfy.model_management.TOTAL_PINNED_MEMORY += size pinned_size[0] += size - _add_to_bucket(module, buckets, size, priority) + _add_to_bucket(module, module_pin, buckets, size, priority) return True diff --git a/comfy/quant_ops.py b/comfy/quant_ops.py index b90bcfd25..53586956a 100644 --- a/comfy/quant_ops.py +++ b/comfy/quant_ops.py @@ -3,6 +3,22 @@ import logging from comfy.cli_args import args + +def _rocm_kitchen_arch_supported(): + """comfy-kitchen's INT8 Triton kernels compile tl.dot to matrix-core instructions. + RDNA3/3.5/4 (gfx11xx/gfx12xx) have WMMA and CDNA (gfx9xx) has MFMA; RDNA1/RDNA2 + (gfx10xx) have neither, so the INT8 path hangs the GPU there. Gates the automatic + ROCm default so those cards stay on the eager fallback (an explicit + --enable-triton-backend still forces it on any arch).""" + try: + arch = torch.cuda.get_device_properties(torch.cuda.current_device()).gcnArchName.split(":")[0] + except Exception: + return False + if arch.startswith(("gfx11", "gfx12")): + return True + return arch in ("gfx908", "gfx90a", "gfx940", "gfx941", "gfx942", "gfx950") + + try: import comfy_kitchen as ck from comfy_kitchen.tensor import ( @@ -10,6 +26,8 @@ try: QuantizedLayout, TensorCoreFP8Layout as _CKFp8Layout, TensorCoreNVFP4Layout as _CKNvfp4Layout, + TensorCoreConvRotW4A4Layout as _CKTensorCoreConvRotW4A4Layout, + TensorWiseINT8Layout as _CKTensorWiseINT8Layout, register_layout_op, register_layout_class, get_layout_class, @@ -23,10 +41,22 @@ try: ck.registry.disable("cuda") logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.") - if args.enable_triton_backend: + # On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated + # comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7 AND a + # matrix-core GPU (RDNA3+ WMMA gfx11xx/gfx12xx, CDNA MFMA gfx9xx). RDNA1/RDNA2 + # (gfx10xx) have no WMMA -> the INT8 tl.dot path hangs the GPU, so they stay eager. + # older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path. + if args.disable_triton_backend: + ck.registry.disable("triton") + elif args.enable_triton_backend: # or (torch.version.hip is not None and _rocm_kitchen_arch_supported()): try: import triton - logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__) + triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2]) + if args.enable_triton_backend or triton_version >= (3, 7): + logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__) + else: + logging.info("Triton %s is too old for the ROCm INT8 path (needs >= 3.7); comfy-kitchen triton backend disabled.", triton.__version__) + ck.registry.disable("triton") except ImportError as e: logging.error(f"Failed to import triton, Error: {e}, the comfy-kitchen triton backend will not be available.") ck.registry.disable("triton") @@ -47,6 +77,12 @@ except ImportError as e: class _CKNvfp4Layout: pass + class _CKTensorWiseINT8Layout: + pass + + class _CKTensorCoreConvRotW4A4Layout: + pass + def register_layout_class(name, cls): pass @@ -174,6 +210,8 @@ class TensorCoreFP8E5M2Layout(_TensorCoreFP8LayoutBase): # Backward compatibility alias - default to E4M3 TensorCoreFP8Layout = TensorCoreFP8E4M3Layout +TensorWiseINT8Layout = _CKTensorWiseINT8Layout +TensorCoreConvRotW4A4Layout = _CKTensorCoreConvRotW4A4Layout # ============================================================================== @@ -184,6 +222,8 @@ register_layout_class("TensorCoreFP8Layout", TensorCoreFP8Layout) register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout) register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout) register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout) +register_layout_class("TensorWiseINT8Layout", _CKTensorWiseINT8Layout) +register_layout_class("TensorCoreConvRotW4A4Layout", _CKTensorCoreConvRotW4A4Layout) if _CK_MXFP8_AVAILABLE: register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout) @@ -200,7 +240,7 @@ QUANT_ALGOS = { }, "nvfp4": { "storage_t": torch.uint8, - "parameters": {"weight_scale", "weight_scale_2", "input_scale"}, + "parameters": {"weight_scale", "weight_scale_2", "input_scale", "pre_quant_scale"}, "comfy_tensor_layout": "TensorCoreNVFP4Layout", "group_size": 16, }, @@ -214,6 +254,20 @@ if _CK_MXFP8_AVAILABLE: "group_size": 32, } +QUANT_ALGOS["int8_tensorwise"] = { + "storage_t": torch.int8, + "parameters": {"weight_scale"}, + "comfy_tensor_layout": "TensorWiseINT8Layout", + "quantize_input": False, +} + +QUANT_ALGOS["convrot_w4a4"] = { + "storage_t": torch.int8, + "parameters": {"weight_scale"}, + "comfy_tensor_layout": "TensorCoreConvRotW4A4Layout", + "quantize_input": False, +} + # ============================================================================== # Re-exports for backward compatibility @@ -226,6 +280,8 @@ __all__ = [ "TensorCoreFP8E4M3Layout", "TensorCoreFP8E5M2Layout", "TensorCoreNVFP4Layout", + "TensorCoreConvRotW4A4Layout", + "TensorWiseINT8Layout", "QUANT_ALGOS", "register_layout_op", ] diff --git a/comfy/samplers.py b/comfy/samplers.py index 25c5a855f..a280f3bb6 100755 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -20,6 +20,7 @@ import comfy.hooks import comfy.context_windows import comfy.multigpu import comfy.utils +from comfy.internal_logging import detail import scipy.stats import numpy @@ -991,10 +992,15 @@ class KSAMPLER(Sampler): noise = model_wrap.inner_model.model_sampling.noise_scaling(sigmas[0], noise, latent_image, self.max_denoise(model_wrap, sigmas)) - k_callback = None total_steps = len(sigmas) - 1 - if callback is not None: - k_callback = lambda x: callback(x["i"], x["denoised"], x["x"], total_steps) + first_step = True + def k_callback(x): + nonlocal first_step + if first_step: + detail("First sampler step: model=%s sampler=%s step=%s total_steps=%s cfg=%s seed=%s sigma=%s sigma_hat=%s latent_shape=%s denoised_shape=%s", model_wrap.model_patcher.model.__class__.__name__, getattr(self.sampler_function, "__name__", "unknown"), x["i"], total_steps, model_wrap.cfg, extra_args.get("seed"), x.get("sigma"), x.get("sigma_hat"), tuple(x["x"].shape), tuple(x["denoised"].shape)) + first_step = False + if callback is not None: + callback(x["i"], x["denoised"], x["x"], total_steps) samples = self.sampler_function(model_k, noise, sigmas, extra_args=extra_args, callback=k_callback, disable=disable_pbar, **self.extra_options) samples = model_wrap.inner_model.model_sampling.inverse_noise_scaling(sigmas[-1], samples) @@ -1270,10 +1276,21 @@ class CFGGuider: return latent_image if latent_image.is_nested: + sampler_shapes = [tuple(x.shape) for x in latent_image.unbind()] latent_image, latent_shapes = comfy.utils.pack_latents(latent_image.unbind()) noise, _ = comfy.utils.pack_latents(noise.unbind()) else: latent_shapes = [latent_image.shape] + sampler_shapes = [tuple(latent_image.shape)] + detail("Sampler: model=%s latent_shapes=%s", self.model_patcher.model.__class__.__name__, sampler_shapes) + + if len(latent_shapes) > 1 and callback is not None: + # samplers run on the flat pack, hand callbacks (previews, x0 output) the nested view + packed_callback = callback + def callback(step, x0, x, total_steps): + x0 = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x0, latent_shapes)) + x = comfy.nested_tensor.NestedTensor(comfy.utils.unpack_latents(x, latent_shapes)) + return packed_callback(step, x0, x, total_steps) if denoise_mask is not None: if denoise_mask.is_nested: diff --git a/comfy/sd.py b/comfy/sd.py index 8e36e7b69..78be9951d 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -16,6 +16,8 @@ import comfy.ldm.cosmos.vae import comfy.ldm.wan.vae import comfy.ldm.wan.vae2_2 import comfy.ldm.hunyuan3d.vae +import comfy.ldm.seedvr.vae +import comfy.ldm.mage_flow.vae import comfy.ldm.triposplat.vae import comfy.ldm.ace.vae.music_dcae_pipeline import comfy.ldm.cogvideo.vae @@ -58,6 +60,8 @@ import comfy.text_encoders.omnigen2 import comfy.text_encoders.qwen_image import comfy.text_encoders.hunyuan_image import comfy.text_encoders.z_image +import comfy.text_encoders.krea2 +import comfy.text_encoders.mage_flow import comfy.text_encoders.ideogram4 import comfy.text_encoders.ovis import comfy.text_encoders.kandinsky5 @@ -68,12 +72,16 @@ import comfy.text_encoders.ace15 import comfy.text_encoders.longcat_image import comfy.text_encoders.qwen35 import comfy.text_encoders.qwen3vl +import comfy.text_encoders.minimax +import comfy.ldm.minimax.vae +import comfy.ldm.minimax.audio_vae import comfy.text_encoders.boogu import comfy.text_encoders.ernie import comfy.text_encoders.gemma4 import comfy.text_encoders.cogvideo import comfy.text_encoders.sa3 import comfy.text_encoders.gpt_oss +import comfy.text_encoders.joyimage import comfy.model_patcher import comfy.lora @@ -467,9 +475,13 @@ class CLIP: def decode(self, token_ids, skip_special_tokens=True): return self.tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens) + def is_dynamic(self): + return self.patcher.is_dynamic() + class VAE: def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None): - if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format + is_seedvr2_vae = "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd + if not is_seedvr2_vae and 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format sd = diffusers_convert.convert_vae_state_dict(sd) if model_management.is_amd(): @@ -496,6 +508,8 @@ class VAE: self.upscale_index_formula = None self.extra_1d_channel = None self.crop_input = True + self.handles_tiling = False + self.format_encoded = None self.audio_sample_rate = 44100 @@ -542,6 +556,33 @@ class VAE: self.first_stage_model = StageC_coder() self.downscale_ratio = 32 self.latent_channels = 16 + elif "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd: # seedvr2 + self.first_stage_model = comfy.ldm.seedvr.vae.VideoAutoencoderKLWrapper() + self.latent_channels = comfy.ldm.seedvr.vae.SEEDVR2_LATENT_CHANNELS + self.latent_dim = 3 + self.disable_offload = True + self.memory_used_decode = lambda shape, dtype: self.first_stage_model.comfy_memory_used_decode(shape) + self.memory_used_encode = lambda shape, dtype: (max(shape[2], 5) * shape[3] * shape[4] * 64) * model_management.dtype_size(dtype) + self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32] + self.handles_tiling = True + self.format_encoded = self.first_stage_model.comfy_format_encoded + self.downscale_ratio = (lambda a: max(0, math.floor((a + 3) / 4)), 8, 8) + self.downscale_index_formula = (4, 8, 8) + self.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8) + self.upscale_index_formula = (4, 8, 8) + self.process_input = lambda image: image * 2.0 - 1.0 + self.crop_input = False + elif "student.dconv_encoder.proj_out.weight" in sd: # Mage-VAE (one-step diffusion codec, Flux2-anchored 128ch/16x latents) + sd = comfy.utils.state_dict_prefix_replace(sd, {"student.dconv_encoder.": "dconv_encoder.", "pipeline.": "decoder_model."}) + # Drop the unused Flux2-VAE anchor encoder carried in the checkpoint. + sd = {k: v for k, v in sd.items() if not k.startswith("decoder_model.y_embedder.encoder.") and not k.startswith("decoder_model.y_embedder.bottleneck.")} + self.first_stage_model = comfy.ldm.mage_flow.vae.MageVAE() + self.latent_channels = 128 + self.downscale_ratio = 16 + self.upscale_ratio = 16 + 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.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} @@ -898,6 +939,50 @@ class VAE: #Force cast it for --disable-dynamic-vram users until there is a true core fix. if not comfy.memory_management.aimdo_enabled: self.disable_offload = True + elif "decoder.transformer_blocks.0.scale1" in sd and "encoder.down.5.block.0.conv1.weight" in sd: # MiniMax H3 video VAE + self.first_stage_model = comfy.ldm.minimax.vae.MiniMaxH3VideoVAE() + self.latent_channels = 24 + self.latent_dim = 3 + # frames 17k+5 <-> latents 5k+2, 16x spatial + self.upscale_ratio = (lambda a: max(1, (a - 2) // 5 * 17 + 5), 16, 16) + self.upscale_index_formula = (4, 16, 16) + self.downscale_ratio = (lambda a: max(1, (a - 5) // 17 * 5 + 2) if a > 1 else 1, 16, 16) + self.downscale_index_formula = (4, 16, 16) + self.working_dtypes = [torch.float16, torch.float32] + # the model tiles internally (256px spatial, 17-frame temporal chunks) + self.handles_tiling = True + 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 + 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 + 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) + self.memory_used_decode = lambda shape, dtype: estimate_decode_memory(self.upscale_ratio[0](shape[2]), shape[3] * self.upscale_ratio[1], shape[4] * self.upscale_ratio[2], dtype) + elif "pre_block.attn.zero_k_bias" in sd: # MiniMax H3 audio VAE (DAC encoder + BigVGAN decoder) + self.first_stage_model = comfy.ldm.minimax.audio_vae.MiniMaxH3AudioVAE() + self.latent_channels = 32 + self.output_channels = 2 + self.pad_channel_value = "replicate" + self.audio_sample_rate = 32000 + self.upscale_ratio = 800 + self.downscale_ratio = 800 + self.latent_dim = 2 # [B, 32, stereo 2, T] + self.process_output = lambda audio: audio + self.process_input = lambda audio: audio + self.working_dtypes = [torch.float32] + # encode gets the waveform shape [B, 2, samples], decode the latent shape [B, 32, 2, T] + def estimate_encode_memory(samples, dtype): + return (900 * samples + 105_000_000) * model_management.dtype_size(dtype) * 1.03 + + def estimate_decode_memory(samples, dtype): + return max(42_000_000, 220 * samples + 20_000_000) * model_management.dtype_size(dtype) * 1.03 + + self.memory_used_encode = lambda shape, dtype: estimate_encode_memory(shape[2], dtype) + self.memory_used_decode = lambda shape, dtype: estimate_decode_memory(shape[-1] * self.upscale_ratio, dtype) elif "gs.base_offset_scale" in sd and "octree.out_proj.weight" in sd: # TripoSplat octree gaussian decoder self.first_stage_model = comfy.ldm.triposplat.vae.OctreeGaussianDecoder() self.latent_channels = 16 @@ -1008,6 +1093,10 @@ class VAE: decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device)) + def _decode_tiled_owned(self, samples, **kwargs): + out = self.first_stage_model.decode_tiled(samples.to(self.vae_dtype).to(self.device), **kwargs) + return self.process_output(out.to(device=self.output_device, dtype=self.vae_output_dtype(), copy=True)) + def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64): steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap) steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap) @@ -1044,6 +1133,25 @@ class VAE: encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype()) return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device) + def _encode_tiled_owned(self, pixel_samples, **kwargs): + x = self.process_input(pixel_samples).to(self.vae_dtype).to(self.device) + out = self.first_stage_model.encode_tiled(x, **kwargs) + return out.to(device=self.output_device, dtype=self.vae_output_dtype()) + + def _owned_tiled_args(self, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): + args = {} + if tile_x is not None: + args["tile_x"] = tile_x + if tile_y is not None: + args["tile_y"] = tile_y + if overlap is not None: + args["overlap"] = overlap + if tile_t is not None: + args["tile_t"] = tile_t + if overlap_t is not None: + args["overlap_t"] = overlap_t + return args + def decode(self, samples_in, vae_options={}): self.throw_exception_if_invalid() pixel_samples = None @@ -1091,11 +1199,19 @@ class VAE: if dims == 1 or self.extra_1d_channel is not None: pixel_samples = self.decode_tiled_1d(samples_in) elif dims == 2: - pixel_samples = self.decode_tiled_(samples_in) + if self.handles_tiling: + tile = 256 // self.spacial_compression_decode() + overlap = tile // 4 + pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) + else: + pixel_samples = self.decode_tiled_(samples_in) elif dims == 3: tile = 256 // self.spacial_compression_decode() overlap = tile // 4 - pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + if self.handles_tiling: + pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap) + else: + pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1) return pixel_samples @@ -1114,7 +1230,9 @@ class VAE: args["overlap"] = overlap with model_management.cuda_device_context(self.device): - if dims == 1 or self.extra_1d_channel is not None: + if self.handles_tiling and dims in (2, 3): + output = self._decode_tiled_owned(samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) + elif dims == 1 or self.extra_1d_channel is not None: args.pop("tile_y") output = self.decode_tiled_1d(samples, **args) elif dims == 2: @@ -1175,12 +1293,17 @@ class VAE: if self.latent_dim == 3: tile = 256 overlap = tile // 4 - samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) + if self.handles_tiling: + samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap) + else: + samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap)) elif self.latent_dim == 1 or self.extra_1d_channel is not None: samples = self.encode_tiled_1d(pixel_samples) else: samples = self.encode_tiled_(pixel_samples) + if self.format_encoded is not None: + samples = self.format_encoded(samples) return samples def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None): @@ -1188,7 +1311,7 @@ class VAE: pixel_samples = self.vae_encode_crop_pixels(pixel_samples) dims = self.latent_dim pixel_samples = pixel_samples.movedim(-1, 1) - if dims == 3: + if dims == 3 and pixel_samples.ndim < 5: if not self.not_video: pixel_samples = pixel_samples.movedim(1, 0).unsqueeze(0) else: @@ -1222,21 +1345,27 @@ class VAE: elif dims == 2: samples = self.encode_tiled_(pixel_samples, **args) elif dims == 3: - if tile_t is not None: - tile_t_latent = max(2, self.downscale_ratio[0](tile_t)) + if self.handles_tiling: + samples = self._encode_tiled_owned(pixel_samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t)) else: - tile_t_latent = 9999 - args["tile_t"] = self.upscale_ratio[0](tile_t_latent) + if tile_t is not None: + tile_t_latent = max(2, self.downscale_ratio[0](tile_t)) + else: + tile_t_latent = 9999 + args["tile_t"] = self.upscale_ratio[0](tile_t_latent) - if overlap_t is None: - args["overlap"] = (1, overlap, overlap) - else: - args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), overlap, overlap) - maximum = pixel_samples.shape[2] - maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum)) + spatial_overlap = overlap if overlap is not None else 64 + if overlap_t is None: + args["overlap"] = (1, spatial_overlap, spatial_overlap) + else: + args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), spatial_overlap, spatial_overlap) + maximum = pixel_samples.shape[2] + maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum)) - samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args) + samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args) + if self.format_encoded is not None: + samples = self.format_encoded(samples) return samples def get_sd(self): @@ -1260,6 +1389,11 @@ class VAE: except: return None + def is_dynamic(self): + # A VAE built from a state dict with no detectable VAE weights returns early + # from __init__ ("No VAE weights detected") before self.patcher is assigned. + patcher = getattr(self, "patcher", None) + return patcher is not None and patcher.is_dynamic() class StyleModel: def __init__(self, model, device="cpu"): @@ -1313,6 +1447,10 @@ class CLIPType(Enum): PIXELDIT = 29 IDEOGRAM4 = 30 BOOGU = 31 + KREA2 = 32 + JOYIMAGE = 33 + MAGE = 34 + MINIMAX = 35 @@ -1368,6 +1506,8 @@ class TEModel(Enum): GPT_OSS_20B = 33 QWEN3VL_4B = 34 QWEN3VL_8B = 35 + GEMMA_4_12B = 36 + QWEN3VL_32B = 37 def detect_te_model(sd): @@ -1397,6 +1537,9 @@ def detect_te_model(sd): if 'model.layers.0.post_feedforward_layernorm.weight' in sd: if 'model.layers.59.self_attn.q_norm.weight' in sd: return TEModel.GEMMA_4_31B + # Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v). + if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd: + return TEModel.GEMMA_4_12B if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd: return TEModel.GEMMA_4_E4B if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd: @@ -1431,6 +1574,9 @@ def detect_te_model(sd): return TEModel.QWEN35_2B if "model.visual.deepstack_merger_list.0.norm.weight" in sd: # DeepStack is unique to Qwen3-VL return TEModel.QWEN3VL_4B if sd["model.visual.merger.linear_fc2.weight"].shape[0] == 2560 else TEModel.QWEN3VL_8B + if "visual.deepstack_merger_list.0.norm.weight" in sd and "model.layers.49.self_attn.q_proj.weight" in sd: + # MiniMax H3 conditioning encoder: Qwen3-VL-32B, truncated to 50 layers + return TEModel.QWEN3VL_32B if "model.layers.0.post_attention_layernorm.weight" in sd: weight = sd['model.layers.0.post_attention_layernorm.weight'] if 'model.layers.0.self_attn.q_norm.weight' in sd: @@ -1552,10 +1698,11 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None) - elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B): + 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}[te_model] + 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) @@ -1638,6 +1785,18 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) clip_target.clip = comfy.text_encoders.boogu.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.boogu.BooguTokenizer + elif clip_type == CLIPType.KREA2 and te_model == TEModel.QWEN3VL_4B: # Krea2: full Qwen3-VL-4B (12-layer tap for conditioning + multimodal generate). + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer + elif clip_type == CLIPType.MAGE and te_model == TEModel.QWEN3VL_4B: # Mage-Flow: full Qwen3-VL-4B, last hidden state, Qwen-Image-style templates. + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.mage_flow.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.mage_flow.MageFlowTokenizer + elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx. + clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."}) + clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.joyimage.JoyImageTokenizer elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused. klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b" clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type) @@ -1647,6 +1806,9 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip qwen3vl_type = {TEModel.QWEN3VL_4B: "qwen3vl_4b", TEModel.QWEN3VL_8B: "qwen3vl_8b"}[te_model] clip_target.clip = comfy.text_encoders.qwen3vl.te(**llama_detect(clip_data), model_type=qwen3vl_type) clip_target.tokenizer = comfy.text_encoders.qwen3vl.tokenizer(model_type=qwen3vl_type) + elif te_model == TEModel.QWEN3VL_32B: + clip_target.clip = comfy.text_encoders.minimax.te(**llama_detect(clip_data)) + clip_target.tokenizer = comfy.text_encoders.minimax.MiniMaxH3Tokenizer elif te_model == TEModel.QWEN3_06B: clip_target.clip = comfy.text_encoders.anima.te(**llama_detect(clip_data)) clip_target.tokenizer = comfy.text_encoders.anima.AnimaTokenizer @@ -1894,7 +2056,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes) else: manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) + model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device) if model_config.clip_vision_prefix is not None: if output_clipvision: @@ -2035,7 +2197,7 @@ def load_diffusion_model_state_dict(sd, model_options={}, metadata=None, disable manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes) else: manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) + model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device) if custom_operations is not None: model_config.custom_operations = custom_operations diff --git a/comfy/sd1_clip.py b/comfy/sd1_clip.py index 897186bba..f0fdf1aa5 100644 --- a/comfy/sd1_clip.py +++ b/comfy/sd1_clip.py @@ -543,18 +543,24 @@ class SDTokenizer: def _try_get_embedding(self, embedding_name:str): ''' Takes a potential embedding name and tries to retrieve it. - Returns a Tuple consisting of the embedding and any leftover string, embedding can be None. + Returns a Tuple consisting of the embedding, the cleaned embedding name, and any leftover string, embedding can be None. ''' split_embed = embedding_name.split() embedding_name = split_embed[0] leftover = ' '.join(split_embed[1:]) + + match = re.search(r'[<\[]', embedding_name) + if match is not None: + leftover = embedding_name[match.start():] + (" " + leftover if leftover else "") + embedding_name = embedding_name[:match.start()] + embed = load_embed(embedding_name, self.embedding_directory, self.embedding_size, self.embedding_key) if embed is None: stripped = embedding_name.strip(',') if len(stripped) < len(embedding_name): embed = load_embed(stripped, self.embedding_directory, self.embedding_size, self.embedding_key) - return (embed, "{} {}".format(embedding_name[len(stripped):], leftover)) - return (embed, leftover) + return (embed, embedding_name, "{} {}".format(embedding_name[len(stripped):], leftover)) + return (embed, embedding_name, leftover) def pad_tokens(self, tokens, amount): if self.pad_left: @@ -585,7 +591,7 @@ class SDTokenizer: tokens = [] for weighted_segment, weight in parsed_weights: to_tokenize = unescape_important(weighted_segment) - split = re.split(' {0}|\n{0}'.format(self.embedding_identifier), to_tokenize) + split = re.split(r'(?<=\s){}'.format(re.escape(self.embedding_identifier)), to_tokenize) to_tokenize = [split[0]] for i in range(1, len(split)): to_tokenize.append("{}{}".format(self.embedding_identifier, split[i])) @@ -595,7 +601,7 @@ class SDTokenizer: # if we find an embedding, deal with the embedding if word.startswith(self.embedding_identifier) and self.embedding_directory is not None: embedding_name = word[len(self.embedding_identifier):].strip('\n') - embed, leftover = self._try_get_embedding(embedding_name) + embed, embedding_name, leftover = self._try_get_embedding(embedding_name) if embed is None: logging.warning(f"warning, embedding:{embedding_name} does not exist, ignoring") else: diff --git a/comfy/supported_models.py b/comfy/supported_models.py index cc05908ee..51b58ed1e 100644 --- a/comfy/supported_models.py +++ b/comfy/supported_models.py @@ -15,6 +15,7 @@ import comfy.text_encoders.flux 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.cosmos import comfy.text_encoders.lumina2 import comfy.text_encoders.wan @@ -26,6 +27,9 @@ import comfy.text_encoders.kandinsky5 import comfy.text_encoders.z_image import comfy.text_encoders.ideogram4 import comfy.text_encoders.boogu +import comfy.text_encoders.krea2 +import comfy.text_encoders.mage_flow +import comfy.text_encoders.joyimage import comfy.text_encoders.anima import comfy.text_encoders.ace15 import comfy.text_encoders.longcat_image @@ -952,6 +956,33 @@ class LTXAV(LTXV): out = model_base.LTXAV(self, device=device) return out +class MiniMaxH3(supported_models_base.BASE): + unet_config = { + "image_model": "minimax_h3", + } + + sampling_settings = { + "shift": 12.0, + } + + unet_extra_config = {} + latent_format = latent_formats.MiniMaxH3AV + + memory_usage_factor = 0.114 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + return model_base.MiniMaxH3(self, device=device) + + def clip_target(self, state_dict={}, prefix=""): + pref = self.text_encoder_key_prefix[0] + detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_32b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.minimax.MiniMaxH3Tokenizer, comfy.text_encoders.minimax.te(**detect)) + class HunyuanVideo(supported_models_base.BASE): unet_config = { "image_model": "hunyuan_video", @@ -1684,6 +1715,40 @@ class Chroma(supported_models_base.BASE): t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.pixart_t5.PixArtTokenizer, comfy.text_encoders.pixart_t5.pixart_te(**t5_detect)) +class SeedVR2(supported_models_base.BASE): + unet_config = { + "image_model": "seedvr2" + } + unet_extra_config = {} + required_keys = { + "{}positive_conditioning", + "{}negative_conditioning", + } + latent_format = comfy.latent_formats.SeedVR2 + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32] + sampling_settings = { + "shift": 1.0, + } + + def set_inference_dtype(self, dtype, manual_cast_dtype, device=None): + if ( + dtype == torch.float16 + and manual_cast_dtype is None + and comfy.model_management.should_use_bf16(device) + ): + manual_cast_dtype = torch.bfloat16 + super().set_inference_dtype(dtype, manual_cast_dtype, device=device) + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.SeedVR2(self, device=device) + return out + + def clip_target(self, state_dict={}): + return None + class ChromaRadiance(Chroma): unet_config = { "image_model": "chroma_radiance", @@ -1818,6 +1883,64 @@ class Ideogram4(supported_models_base.BASE): hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_8b.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.ideogram4.Ideogram4Tokenizer, comfy.text_encoders.ideogram4.te(**hunyuan_detect)) + +class Krea2(supported_models_base.BASE): + unet_config = { + "image_model": "krea2", + } + + sampling_settings = { + "multiplier": 1.0, + "shift": 1.15, + } + + memory_usage_factor = 2.2 + + latent_format = latent_formats.Wan21 + + supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.Krea2(self, device=device) + return out + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.krea2.Krea2Tokenizer, comfy.text_encoders.krea2.te(**hunyuan_detect)) + +class MageFlow(supported_models_base.BASE): + unet_config = { + "image_model": "mage_flow", + } + + sampling_settings = { + "multiplier": 1.0, + "shift": 6.0, + } + + memory_usage_factor = 6.5 + + unet_extra_config = {} + latent_format = latent_formats.Flux2 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + out = model_base.MageFlow(self, device=device) + return out + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl_4b.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.mage_flow.MageFlowTokenizer, comfy.text_encoders.mage_flow.te(**hunyuan_detect)) + class QwenImage(supported_models_base.BASE): unet_config = { "image_model": "qwen_image", @@ -1847,6 +1970,38 @@ class QwenImage(supported_models_base.BASE): hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref)) return supported_models_base.ClipTarget(comfy.text_encoders.qwen_image.QwenImageTokenizer, comfy.text_encoders.qwen_image.te(**hunyuan_detect)) +class JoyImage(supported_models_base.BASE): + unet_config = { + "image_model": "joyimage", + } + + sampling_settings = { + "multiplier": 1000, + "shift": 1.5, + } + + memory_usage_factor = 1.8 + + unet_extra_config = { + "theta": 10000, + "rope_dim_list": [16, 56, 56], + } + + latent_format = latent_formats.Wan21 + + supported_inference_dtypes = [torch.bfloat16, torch.float32] + + vae_key_prefix = ["vae."] + text_encoder_key_prefix = ["text_encoders."] + + def get_model(self, state_dict, prefix="", device=None): + return model_base.JoyImage(self, device=device) + + def clip_target(self, state_dict={}): + pref = self.text_encoder_key_prefix[0] + qwen3vl_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl.transformer.".format(pref)) + return supported_models_base.ClipTarget(comfy.text_encoders.joyimage.JoyImageTokenizer, comfy.text_encoders.joyimage.te(**qwen3vl_detect)) + class HunyuanImage21(HunyuanVideo): unet_config = { "image_model": "hunyuan_video", @@ -2280,6 +2435,7 @@ models = [ GenmoMochi, LTXV, LTXAV, + MiniMaxH3, HunyuanVideo15_SR_Distilled, HunyuanVideo15, HunyuanImage21Refiner, @@ -2318,13 +2474,17 @@ models = [ HiDream, HiDreamO1, Chroma, + SeedVR2, ChromaRadiance, ACEStep, ACEStep15, Omnigen2, Boogu, + MageFlow, QwenImage, + JoyImage, Ideogram4, + Krea2, Flux2, Lens, Kandinsky5Image, diff --git a/comfy/supported_models_base.py b/comfy/supported_models_base.py index 0e7a829ba..e3a8e131f 100644 --- a/comfy/supported_models_base.py +++ b/comfy/supported_models_base.py @@ -54,13 +54,13 @@ class BASE: optimizations = {"fp8": False} @classmethod - def matches(s, unet_config, state_dict=None): + def matches(s, unet_config, state_dict=None, unet_key_prefix=""): for k in s.unet_config: if k not in unet_config or s.unet_config[k] != unet_config[k]: return False if state_dict is not None: for k in s.required_keys: - if k not in state_dict: + if k.format(unet_key_prefix) not in state_dict: return False return True @@ -115,7 +115,7 @@ class BASE: replace_prefix = {"": self.vae_key_prefix[0]} return utils.state_dict_prefix_replace(state_dict, replace_prefix) - def set_inference_dtype(self, dtype, manual_cast_dtype): + def set_inference_dtype(self, dtype, manual_cast_dtype, device=None): self.unet_config['dtype'] = dtype self.manual_cast_dtype = manual_cast_dtype diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index f050061ed..5163c1676 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1,11 +1,15 @@ import torch import torch.nn as nn +import torchaudio.functional as AF +import torchvision.transforms.functional as TVF import numpy as np +from tokenizers import Tokenizer from dataclasses import dataclass import math from comfy import sd1_clip import comfy.model_management +import comfy.ops from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.rmsnorm import rms_norm from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding @@ -21,6 +25,10 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3} GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5} +# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space. +GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6} +GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6} + @dataclass class Gemma4Config: vocab_size: int = 262144 @@ -35,6 +43,9 @@ class Gemma4Config: transformer_type: str = "gemma4" head_dim = 256 global_head_dim = 512 + num_global_key_value_heads = None + attention_k_eq_v = False + vision_bidirectional = False rms_norm_add = False mlp_activation = "gelu_pytorch_tanh" qkv_bias = False @@ -51,6 +62,7 @@ class Gemma4Config: num_kv_shared_layers: int = 18 use_double_wide_mlp: bool = False stop_tokens = [1, 50, 106] + suppress_tokens = [] vision_config = GEMMA4_VISION_CONFIG audio_config = GEMMA4_AUDIO_CONFIG mm_tokens_per_image = 280 @@ -72,12 +84,30 @@ class Gemma4_31B_Config(Gemma4Config): num_hidden_layers: int = 60 num_attention_heads: int = 32 num_key_value_heads: int = 16 + vision_bidirectional = True sliding_attention = [1024, 1024, 1024, 1024, 1024, False] hidden_size_per_layer_input: int = 0 num_kv_shared_layers: int = 0 audio_config = None vision_config = GEMMA4_VISION_31B_CONFIG +@dataclass +class Gemma4_12B_Config(Gemma4Config): + hidden_size: int = 3840 + intermediate_size: int = 15360 + num_hidden_layers: int = 48 + num_attention_heads: int = 16 + num_key_value_heads: int = 8 + num_global_key_value_heads = 1 + attention_k_eq_v = True + vision_bidirectional = True + sliding_attention = [1024, 1024, 1024, 1024, 1024, False] + hidden_size_per_layer_input: int = 0 + num_kv_shared_layers: int = 0 + audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG + vision_config = GEMMA4_UNIFIED_VISION_CONFIG + suppress_tokens = [258883, 258882] + # unfused RoPE as addcmul_ RoPE diverges from reference code def _apply_rotary_pos_emb(x, freqs_cis): @@ -89,17 +119,18 @@ def _apply_rotary_pos_emb(x, freqs_cis): return out class Gemma4Attention(nn.Module): - def __init__(self, config, head_dim, device=None, dtype=None, ops=None): + def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None): super().__init__() self.num_heads = config.num_attention_heads - self.num_kv_heads = config.num_key_value_heads + self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads self.hidden_size = config.hidden_size self.head_dim = head_dim self.inner_size = self.num_heads * head_dim self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype) self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) - self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) + # k_eq_v: V reuses the K projection (no separate v_proj weight) + self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype) self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype) self.q_norm = None @@ -133,7 +164,10 @@ class Gemma4Attention(nn.Module): shareable_kv = None else: xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) - xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) + if self.v_proj is not None: + xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim) + else: + xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE) if self.k_norm is not None: xk = self.k_norm(xk) xv = rms_norm(xv) @@ -186,7 +220,10 @@ class TransformerBlockGemma4(nn.Module): head_dim = config.head_dim if self.sliding_attention else config.global_head_dim - self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops) + # k_eq_v only on global layers, which then use num_global_key_value_heads + k_eq_v = config.attention_k_eq_v and not self.sliding_attention + num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads + self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops) num_kv_shared = config.num_kv_shared_layers first_kv_shared = config.num_hidden_layers - num_kv_shared @@ -203,9 +240,9 @@ class TransformerBlockGemma4(nn.Module): self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype) self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype) self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype) - self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype)) - else: - self.layer_scalar = None + + # layer_scalar exists on every gemma4 variant, independent of per-layer input + 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): sliding_window = None @@ -244,8 +281,7 @@ class TransformerBlockGemma4(nn.Module): x = self.post_per_layer_input_norm(x) x = residual + x - if self.layer_scalar is not None: - x = x * self.layer_scalar + x = x * comfy.ops.cast_to_input(self.layer_scalar, x) return x, present_key_value, shareable_kv @@ -334,6 +370,19 @@ class Gemma4Transformer(nn.Module): causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val) mask = mask + causal_mask if mask is not None else causal_mask + # Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal). + if self.config.vision_bidirectional and past_len == 0 and embeds_info: + block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device) + group = 0 + for info in embeds_info: + if info.get("type") == "image": + start = info["index"] + block_ids[start:start + info["size"]] = group + group += 1 + if group > 0: + same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0) + mask = mask.masked_fill(same_block, 0.0) + # Per-layer inputs per_layer_inputs = None if self.hidden_size_per_layer_input: @@ -354,8 +403,24 @@ class Gemma4Transformer(nn.Module): shared_global_kv = None # KV from last non-shared global layer intermediate = None + all_intermediate = None + only_layers = None + if intermediate_output is not None: + if isinstance(intermediate_output, list): + all_intermediate = [] + only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output} + elif intermediate_output == "all": + all_intermediate = [] + intermediate_output = None + elif intermediate_output < 0: + intermediate_output = len(self.layers) + intermediate_output + next_key_values = [] for i, layer in enumerate(self.layers): + if all_intermediate is not None: + if only_layers is None or (i in only_layers): + all_intermediate.append(x.unsqueeze(1).clone()) + past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None layer_kwargs = {} @@ -385,7 +450,18 @@ class Gemma4Transformer(nn.Module): if self.norm is not None: x = self.norm(x) - if len(next_key_values) > 0: + if all_intermediate is not None: + if only_layers is None or (len(self.layers) in only_layers): + all_intermediate.append(x.unsqueeze(1).clone()) + if len(all_intermediate) > 0: + intermediate = torch.cat(all_intermediate, dim=1) + + if intermediate is not None and final_layer_norm_intermediate and self.norm is not None: + intermediate = self.norm(intermediate) + + # Only hand back the KV cache when caching was actually requested; SDClipModel reads + # outputs[2] as the pooled output. + if past_key_values is not None and len(next_key_values) > 0: return x, intermediate, next_key_values return x, intermediate @@ -404,6 +480,8 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module): cap = self.model.config.final_logit_softcapping if cap: logits = cap * torch.tanh(logits / cap) + if self.model.config.suppress_tokens: + logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min return logits def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): @@ -441,6 +519,28 @@ class Gemma4AudioMixin: return None, None +class Gemma4UnifiedBase(Gemma4Base): + """Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space.""" + def _init_model(self, config, dtype, device, operations): + self.num_layers = config.num_hidden_layers + self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations) + self.dtype = dtype + self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations) + self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations) + self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations) + + def preprocess_embed(self, embed, device): + if embed["type"] == "image": + pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1] + patches, positions = self.vision_model.patchify(pixels) + vision_out = self.vision_model(patches, positions) + return self.multi_modal_projector(vision_out), None + if embed["type"] == "audio": + audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token] + return self.audio_projector(audio), None + return None, None + + # Vision Encoder def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None): @@ -713,6 +813,73 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector): super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops) +# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space. + +def _patches_merge(patches, positions_xy, length): + patch_size = math.isqrt(patches.shape[-1] // 3) + k = math.isqrt(patches.shape[-2] // length) + batch = patches.shape[:-2] + + max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1 + kidx = torch.div(positions_xy, k, rounding_mode="floor") + rem = torch.remainder(positions_xy, k) + order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1] + perm = order.long().argsort(dim=-1) + + merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches)) + merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3) + merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3) + + pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy)) + pad = (positions_xy == -1).all(dim=-1, keepdim=True) + pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2) + pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0] + return merged, pos + + +class Gemma4UnifiedVisionEmbedder(nn.Module): + """Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector.""" + def __init__(self, config, device=None, dtype=None, ops=None): + super().__init__() + self.patch_size = config["patch_size"] + self.pooling_kernel_size = config["pooling_kernel_size"] + patch_dim = config["model_patch_size"] ** 2 * 3 + mm_embed_dim = config["mm_embed_dim"] + self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype) + self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype) + self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype) + self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype)) + self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype) + + def patchify(self, pixels): + """pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2].""" + ps, k = self.patch_size, self.pooling_kernel_size + out_patches, out_positions = [], [] + for img in pixels: + ph, pw = img.shape[-2] // ps, img.shape[-1] // ps + teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1) + grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy") + tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2) + n_model = teacher.shape[0] // (k * k) + mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model) + out_patches.append(mp.squeeze(0)) + out_positions.append(mpos.squeeze(0)) + return torch.stack(out_patches), torch.stack(out_positions) + + def forward(self, pixel_values, image_position_ids): + x = self.patch_ln1(pixel_values) + x = self.patch_dense(x) + x = self.patch_ln2(x) + + clamped = image_position_ids.clamp(min=0).long() + valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1) + axes = torch.arange(2, device=image_position_ids.device) + pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype) + pos_embs = (pos[clamped, axes] * valid).sum(-2) + x = x + pos_embs + return self.pos_norm(x) + + # Audio Encoder class Gemma4AudioConvSubsampler(nn.Module): @@ -990,6 +1157,30 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector): # Tokenizer and Wrappers +def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size): + target_px = max_patches * patch_size ** 2 + factor = math.sqrt(target_px / (height * width)) + side_mult = pooling_kernel_size * patch_size + target_height = math.floor(factor * height / side_mult) * side_mult + target_width = math.floor(factor * width / side_mult) * side_mult + + if target_height == 0 and target_width == 0: + raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.") + + max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult + if target_height == 0: + target_height = side_mult + target_width = min(math.floor(width / height) * side_mult, max_side_length) + elif target_width == 0: + target_width = side_mult + target_height = min(math.floor(height / width) * side_mult, max_side_length) + + if target_height * target_width > target_px: + raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.") + + return target_height, target_width + + class Gemma4_Tokenizer(): tokenizer_json_data = None @@ -998,25 +1189,35 @@ class Gemma4_Tokenizer(): return {"tokenizer_json": self.tokenizer_json_data} return {} - def _extract_mel_spectrogram(self, waveform, sample_rate): - """Extract 128-bin log mel spectrogram. - Uses numpy for FFT/matmul/log to produce bit-identical results with reference code. - """ - # Mix to mono first, then resample to 16kHz + def _audio_token_count(self, num_samples): + # Default (E2B/E4B): mel frames after two stride-2 conv subsamples. + _fl = 320 # int(round(16000 * 20.0 / 1000.0)) + _hl = 160 # int(round(16000 * 10.0 / 1000.0)) + _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1 + _t = _nmel + for _ in range(2): + _t = (_t + 2 - 3) // 2 + 1 + return min(_t, 750) + + @staticmethod + def _resample_16k(waveform, sample_rate): + """Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers + load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio.""" if waveform.dim() > 1 and waveform.shape[0] > 1: waveform = waveform.mean(dim=0, keepdim=True) if waveform.dim() == 1: waveform = waveform.unsqueeze(0) - audio = waveform.squeeze(0).float().numpy() + audio = waveform.float() if sample_rate != 16000: - # Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match) - from scipy.signal import resample_poly, firwin - from math import gcd - g = gcd(sample_rate, 16000) - up, down = 16000 // g, sample_rate // g - L = max(up, down) - h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5)) - audio = resample_poly(audio, up, down, window=h).astype(np.float32) + audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser", + lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614) + return audio.squeeze(0).contiguous() + + def _extract_audio_features(self, waveform, sample_rate): + """Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder. + Uses numpy for FFT/matmul/log to produce bit-identical results with reference code. + """ + audio = self._resample_16k(waveform, sample_rate).numpy() n = len(audio) # Pad to multiple of 128, build sample-level mask @@ -1064,8 +1265,8 @@ class Gemma4_Tokenizer(): if audio is not None: waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000 - mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate) - audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T]) + feat, feat_mask = self._extract_audio_features(waveform, sample_rate) + audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T]) # Process image/video frames is_video = video is not None @@ -1088,15 +1289,10 @@ class Gemma4_Tokenizer(): h, w = samples.shape[2], samples.shape[3] patch_size = 16 pooling_k = 3 - max_soft_tokens = 70 if is_video else 280 # video uses smaller token budget per frame + max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280) max_patches = max_soft_tokens * pooling_k * pooling_k - target_px = max_patches * patch_size * patch_size - factor = (target_px / (h * w)) ** 0.5 - side_mult = pooling_k * patch_size - target_h = max(int(factor * h // side_mult) * side_mult, side_mult) - target_w = max(int(factor * w // side_mult) * side_mult, side_mult) + target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k) - import torchvision.transforms.functional as TVF for i in range(num_frames): # rescaling to match reference code s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8 @@ -1115,7 +1311,7 @@ class Gemma4_Tokenizer(): llama_text = llama_template.format(text) else: # Build template from modalities present - system = "<|turn>system\n<|think|>\n" if thinking else "" + system = "<|turn>system\n<|think|>\n\n" if thinking else "" media = "" if len(images) > 0: if is_video: @@ -1135,15 +1331,11 @@ class Gemma4_Tokenizer(): if len(audio_features) > 0: # Compute audio token count (always at 16kHz) num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1] - _fl = 320 # int(round(16000 * 20.0 / 1000.0)) - _hl = 160 # int(round(16000 * 10.0 / 1000.0)) - _nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1 - _t = _nmel - for _ in range(2): - _t = (_t + 2 - 3) // 2 + 1 - n_audio_tokens = min(_t, 750) + n_audio_tokens = self._audio_token_count(num_samples) media += "<|audio>" + "<|audio|>" * n_audio_tokens + "" - llama_text = f"{system}<|turn>user\n{media}{text}\n<|turn>model\n" + # Non-thinking mode primes an empty thought channel so the model answers directly. + model_open = "" if thinking else "<|channel>thought\n" + 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) @@ -1178,7 +1370,6 @@ class Gemma4_Tokenizer(): class _Gemma4Tokenizer: """Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)""" def __init__(self, tokenizer_json_bytes=None, **kwargs): - from tokenizers import Tokenizer if isinstance(tokenizer_json_bytes, torch.Tensor): tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist()) self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8")) @@ -1224,6 +1415,30 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class) +class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer): + """Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram.""" + embedding_size = 3840 + + def _extract_audio_features(self, waveform, sample_rate): + audio = self._resample_16k(waveform, sample_rate) + spt = 640 # audio_samples_per_token (40ms at 16kHz) + pad = (-audio.shape[0]) % spt + if pad: + audio = torch.nn.functional.pad(audio, (0, pad)) + num_tokens = audio.shape[0] // spt + feats = audio[:num_tokens * spt].reshape(num_tokens, spt) + feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s) + mask = torch.ones(feats.shape[0], dtype=torch.bool) + return feats, mask + + def _audio_token_count(self, num_samples): + return min((num_samples + 639) // 640, 750) + + +class Gemma4UnifiedTokenizer(Gemma4Tokenizer): + tokenizer_class = Gemma4UnifiedSDTokenizer + + # Model wrappers class Gemma4Model(sd1_clip.SDClipModel): model_class = None @@ -1256,7 +1471,7 @@ class Gemma4Model(sd1_clip.SDClipModel): expanded_idx += 1 initial_token_ids = [ids] input_ids = torch.tensor(initial_token_ids, device=self.execution_device) - return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids) + 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_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None): @@ -1296,3 +1511,11 @@ def _make_variant(config_cls): Gemma4_E4B = _make_variant(Gemma4Config) Gemma4_E2B = _make_variant(Gemma4_E2B_Config) Gemma4_31B = _make_variant(Gemma4_31B_Config) + + +# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant). +class Gemma4_12B(Gemma4UnifiedBase): + def __init__(self, config_dict, dtype, device, operations): + super().__init__() + self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations) +Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer diff --git a/comfy/text_encoders/gpt_oss.py b/comfy/text_encoders/gpt_oss.py index d596ef9a0..066796b6a 100644 --- a/comfy/text_encoders/gpt_oss.py +++ b/comfy/text_encoders/gpt_oss.py @@ -12,7 +12,7 @@ import torch.nn.functional as F import comfy.ops from comfy import sd1_clip -from comfy.ldm.modules.attention import TORCH_HAS_GQA, optimized_attention_for_device +from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.text_encoders.llama import RMSNorm, apply_rope @@ -110,10 +110,6 @@ def _attention_with_sinks(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, sin putting the sink logit in the mask at that column. """ - if num_kv_groups > 1 and not TORCH_HAS_GQA: - k = k.repeat_interleave(num_kv_groups, dim=1) - v = v.repeat_interleave(num_kv_groups, dim=1) - B, _, S_q, D = q.shape H_kv = k.shape[1] S_kv = k.shape[-2] diff --git a/comfy/text_encoders/joyimage.py b/comfy/text_encoders/joyimage.py new file mode 100644 index 000000000..143c44250 --- /dev/null +++ b/comfy/text_encoders/joyimage.py @@ -0,0 +1,97 @@ +import torch + +from comfy import sd1_clip +import comfy.text_encoders.qwen_vl +from comfy.text_encoders.qwen3vl import Qwen3VL, Qwen3VLTokenizer + +JOYIMAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>" +JOYIMAGE_TEMPLATE_TEXT = ( + "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" + "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" +) +JOYIMAGE_TEMPLATE_IMAGE = ( + "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " + "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" + f"<|im_start|>user\n{JOYIMAGE_VISION_BLOCK}{{}}<|im_end|>\n<|im_start|>assistant\n" +) +# The DiT was trained without the leading system-prompt tokens. +JOYIMAGE_DROP_IDX = 34 +PAD_TOKEN = 151643 + + +class Qwen3VL8B_JoyImage(Qwen3VL): + model_type = "qwen3vl_8b" + + def preprocess_embed(self, embed, device): + if embed["type"] == "image": + image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images( + embed["data"], min_pixels=65536, max_pixels=16777216, patch_size=16, + image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5], + interpolation="bicubic", + ) + merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid) + return merged, {"grid": grid, "deepstack": deepstack} + return None, None + + +class JoyImageTokenizer(Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__( + embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, + model_type="qwen3vl_8b", + ) + self.llama_template = JOYIMAGE_TEMPLATE_TEXT + self.llama_template_images = JOYIMAGE_TEMPLATE_IMAGE + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=None, **kwargs): + kwargs.pop("thinking", None) + return super().tokenize_with_weights( + text, return_word_ids=return_word_ids, llama_template=llama_template, + images=images or [], thinking=True, **kwargs, + ) + + +class _JoyImageClipModel(sd1_clip.SDClipModel): + def __init__(self, device="cpu", layer="hidden", layer_idx=-1, dtype=None, + attention_mask=True, model_options={}): + super().__init__( + device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={}, + # JoyImage conditions on the pre-final-norm output of the last decoder layer. + dtype=dtype, special_tokens={"pad": PAD_TOKEN}, layer_norm_hidden_state=False, + model_class=Qwen3VL8B_JoyImage, enable_attention_masks=attention_mask, + return_attention_masks=attention_mask, model_options=model_options, + ) + + +class JoyImageTEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + super().__init__( + device=device, dtype=dtype, name="qwen3vl_8b", + clip_model=_JoyImageClipModel, model_options=model_options, + ) + + def encode_token_weights(self, token_weight_pairs): + out, pooled, extra = super().encode_token_weights(token_weight_pairs) + if out.shape[1] <= JOYIMAGE_DROP_IDX: + raise ValueError( + f"JoyImageTEModel: encoded sequence length {out.shape[1]} is shorter " + f"than drop_idx={JOYIMAGE_DROP_IDX}; the prompt did not include the " + f"template prefix." + ) + out = out[:, JOYIMAGE_DROP_IDX:] + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, JOYIMAGE_DROP_IDX:] + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class JoyImageTEModel_(JoyImageTEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + if dtype_llama is not None: + dtype = dtype_llama + super().__init__(device=device, dtype=dtype, model_options=model_options) + return JoyImageTEModel_ diff --git a/comfy/text_encoders/krea2.py b/comfy/text_encoders/krea2.py new file mode 100644 index 000000000..408a03566 --- /dev/null +++ b/comfy/text_encoders/krea2.py @@ -0,0 +1,84 @@ +"""Krea 2 (K2) text encoder: Qwen3-VL-4B, 12-layer tap. + +K2 conditions on a stack of hidden states from 12 layers of Qwen3-VL-4B +(reference taps ``hidden_states[2,5,8,...,35]``), kept as a ``(B, 12, seq, 2560)`` tensor and +consumed by the DiT's internal ``txtfusion`` adapter. Comfy carries conditioning as a 3D tensor, +so the 12-layer stack is flattened to ``(B, seq, 12*2560)`` here and unpacked inside the model. +""" + +import numbers + +import torch + +import comfy.text_encoders.qwen3vl +from comfy import sd1_clip + +# tap k == hidden_states[k] (no offset). +KREA2_TAP_LAYERS = [2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35] + +# Identical system template to Qwen-Image; Krea2 strips the system+user-opening prefix. +KREA2_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + + +class Krea2Tokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b") + self.llama_template = KREA2_TEMPLATE # conditioning template; image text-gen uses qwen3vl's default image template. + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs): + # Krea2 conditions on the no-think template; thinking=True drops the empty block qwen3vl adds. + return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs) + + +class Krea2Qwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel): + def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}): + super().__init__(device=device, layer=KREA2_TAP_LAYERS, layer_idx=None, dtype=dtype, + attention_mask=attention_mask, model_options=model_options, model_type="qwen3vl_4b") + + +class Krea2TEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=Krea2Qwen3VLClipModel, model_options=model_options) + + def encode_token_weights(self, token_weight_pairs, template_end=-1): + out, pooled, extra = super().encode_token_weights(token_weight_pairs) # out: (B, 12, seq, 2560) + tok_pairs = token_weight_pairs["qwen3vl_4b"][0] + + # Strip the system + user-opening prefix + count_im_start = 0 + if template_end == -1: + for i, v in enumerate(tok_pairs): + elem = v[0] + if not torch.is_tensor(elem) and isinstance(elem, numbers.Integral): + if elem == 151644 and count_im_start < 2: + template_end = i + count_im_start += 1 + if out.shape[2] > (template_end + 3): + if tok_pairs[template_end + 1][0] == 872: # "user" + if tok_pairs[template_end + 2][0] == 198: # "\n" + template_end += 3 + + out = out[:, :, template_end:] + + b, n, seq, h = out.shape + # Flatten the 12-layer axis into the feature dim: (B, seq, 12*2560). Unpacked in the model. + out = out.permute(0, 2, 1, 3).reshape(b, seq, n * h) + + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, template_end:] + if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]): + extra.pop("attention_mask") + + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class Krea2TEModel_(Krea2TEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + if dtype_llama is not None: + dtype = dtype_llama + super().__init__(device=device, dtype=dtype, model_options=model_options) + return Krea2TEModel_ diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index e9f38a9a2..f5c5597ef 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -264,6 +264,17 @@ class Qwen3VL_4BConfig(Qwen3VL_8BConfig): intermediate_size: int = 9728 lm_head: bool = False # 4B ties word embeddings +@dataclass +class Qwen3VL_32BConfig(Qwen3VL_8BConfig): + # MiniMax H3 conditioning checkpoint: truncated to the first 50 of 64 layers, + # consumed as the unnormalized hidden state after layer 50 (no final norm, no lm_head) + hidden_size: int = 5120 + intermediate_size: int = 25600 + num_hidden_layers: int = 50 + num_attention_heads: int = 64 + lm_head: bool = False + final_norm: bool = False + @dataclass class Ovis25_2BConfig: vocab_size: int = 151936 @@ -550,10 +561,8 @@ class Attention(nn.Module): xv = xv[:, :, -sliding_window:] attention_mask = attention_mask[..., -sliding_window:] if attention_mask is not None else None - xk = xk.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - xv = xv.repeat_interleave(self.num_heads // self.num_kv_heads, dim=1) - - output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True) + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True, **gqa_kwargs) return self.o_proj(output), present_key_value class MLP(nn.Module): @@ -878,7 +887,7 @@ class BaseGenerate: torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0)) return past_key_values - def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None): + 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 if stop_tokens is None: @@ -913,7 +922,7 @@ class BaseGenerate: if step == 0 and deepstack_embeds is not None: extra["deepstack_embeds"] = deepstack_embeds extra["visual_pos_masks"] = visual_pos_masks - x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra) + x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None)) logits = self.logits(x)[:, -1] next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty) token_id = next_token[0].item() @@ -937,22 +946,41 @@ class BaseGenerate: return torch.argmax(logits, dim=-1, keepdim=True) # Sampling mode - if repetition_penalty != 1.0: - for i in range(logits.shape[0]): - for token_id in set(token_history): - logits[i, token_id] *= repetition_penalty if logits[i, token_id] < 0 else 1/repetition_penalty - - if presence_penalty is not None and presence_penalty != 0.0: - for i in range(logits.shape[0]): - for token_id in set(token_history): - logits[i, token_id] -= presence_penalty + if len(token_history) > 0 and (repetition_penalty != 1.0 or (presence_penalty is not None and presence_penalty != 0.0)): + token_ids = torch.tensor(list(set(token_history)), device=logits.device) + token_logits = logits[:, token_ids] + if repetition_penalty != 1.0: + token_logits = torch.where(token_logits < 0, token_logits * repetition_penalty, token_logits / repetition_penalty) + if presence_penalty is not None and presence_penalty != 0.0: + token_logits = token_logits - presence_penalty + logits[:, token_ids] = token_logits if temperature != 1.0: logits = logits / temperature if top_k > 0: - indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None] - logits[indices_to_remove] = torch.finfo(logits.dtype).min + top_k = min(top_k, logits.shape[-1]) + logits, top_indices = torch.topk(logits, top_k) + + if min_p > 0.0: + probs_before_filter = torch.nn.functional.softmax(logits, dim=-1) + top_probs, _ = probs_before_filter.max(dim=-1, keepdim=True) + min_threshold = min_p * top_probs + indices_to_remove = probs_before_filter < min_threshold + logits[indices_to_remove] = torch.finfo(logits.dtype).min + + if top_p < 1.0: + sorted_logits, sorted_indices = torch.sort(logits, descending=True) + cumulative_probs = torch.cumsum(torch.nn.functional.softmax(sorted_logits, dim=-1), dim=-1) + sorted_indices_to_remove = cumulative_probs > top_p + sorted_indices_to_remove[..., 0] = False + indices_to_remove = torch.zeros_like(logits, dtype=torch.bool) + indices_to_remove.scatter_(1, sorted_indices, sorted_indices_to_remove) + logits[indices_to_remove] = torch.finfo(logits.dtype).min + + probs = torch.nn.functional.softmax(logits, dim=-1) + next_token = torch.multinomial(probs, num_samples=1, generator=generator) + return top_indices.gather(1, next_token) if min_p > 0.0: probs_before_filter = torch.nn.functional.softmax(logits, dim=-1) diff --git a/comfy/text_encoders/mage_flow.py b/comfy/text_encoders/mage_flow.py new file mode 100644 index 000000000..6542ad315 --- /dev/null +++ b/comfy/text_encoders/mage_flow.py @@ -0,0 +1,94 @@ +"""Mage-Flow text encoder: Qwen3-VL-4B, last hidden state (2560-dim). + +Mage-Flow conditions on the final hidden state of Qwen3-VL-4B with the leading +system + user-opening template tokens stripped (reference start_idx 34 for t2i, +64 for edit). The t2i template is identical to Qwen-Image's; the edit template +uses the same system prompt as Qwen-Image-Edit with "Image N: " reference +prefixes and no block. +""" + +import numbers + +import torch + +import comfy.text_encoders.qwen3vl +from comfy import sd1_clip + +MAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>" + +MAGE_T2I_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" +MAGE_EDIT_TEMPLATE = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + + +class MageFlowTokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): + def __init__(self, embedding_directory=None, tokenizer_data={}): + super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b") + self.llama_template = MAGE_T2I_TEMPLATE + self.llama_template_images = MAGE_EDIT_TEMPLATE + + def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs): + image = kwargs.get("image", None) + if image is not None and len(images) == 0: + images = [image[i:i + 1] for i in range(image.shape[0])] + if llama_template is None: + if len(images) > 0: + # Training-time multi-reference body: "Image 1: Image 2: ...{instruction}" + prefix = "".join("Image {}: {}".format(j + 1, MAGE_VISION_BLOCK) for j in range(len(images))) + llama_template = self.llama_template_images.replace("{}", prefix + "{}", 1) + else: + llama_template = self.llama_template + # thinking=True: Mage templates end at "<|im_start|>assistant\n" with no block. + return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs) + + +class MageFlowQwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel): + def __init__(self, device="cpu", dtype=None, attention_mask=True, model_options={}, model_type="qwen3vl_4b"): + super().__init__(device=device, dtype=dtype, attention_mask=attention_mask, model_options=model_options, model_type=model_type) + # apply the final RMSNorm to the tapped last layer (HF last_hidden_state) + self.layer_norm_hidden_state = True + + +class MageFlowTEModel(sd1_clip.SD1ClipModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + clip_model = lambda **kw: MageFlowQwen3VLClipModel(**kw, model_type="qwen3vl_4b") # noqa: E731 + super().__init__(device=device, dtype=dtype, name="qwen3vl_4b", clip_model=clip_model, model_options=model_options) + + def encode_token_weights(self, token_weight_pairs, template_end=-1): + # Strip the system + user-opening prefix (reference drop_idx: 34 t2i / 64 edit). + out, pooled, extra = super().encode_token_weights(token_weight_pairs) + tok_pairs = token_weight_pairs["qwen3vl_4b"][0] + count_im_start = 0 + if template_end == -1: + for i, v in enumerate(tok_pairs): + elem = v[0] + if not torch.is_tensor(elem): + if isinstance(elem, numbers.Integral): + if elem == 151644 and count_im_start < 2: # <|im_start|> + template_end = i + count_im_start += 1 + + if out.shape[1] > (template_end + 3): + if tok_pairs[template_end + 1][0] == 872: # "user" + if tok_pairs[template_end + 2][0] == 198: # "\n" + template_end += 3 + + out = out[:, template_end:] + + if "attention_mask" in extra: + extra["attention_mask"] = extra["attention_mask"][:, template_end:] + if extra["attention_mask"].sum() == torch.numel(extra["attention_mask"]): + extra.pop("attention_mask") # attention mask is useless if no masked elements + + return out, pooled, extra + + +def te(dtype_llama=None, llama_quantization_metadata=None): + class MageFlowTEModel_(MageFlowTEModel): + def __init__(self, device="cpu", dtype=None, model_options={}): + if dtype_llama is not None: + dtype = dtype_llama + if llama_quantization_metadata is not None: + model_options = model_options.copy() + model_options["quantization_metadata"] = llama_quantization_metadata + super().__init__(device=device, dtype=dtype, model_options=model_options) + return MageFlowTEModel_ diff --git a/comfy/text_encoders/minimax.py b/comfy/text_encoders/minimax.py new file mode 100644 index 000000000..c2dc47f7f --- /dev/null +++ b/comfy/text_encoders/minimax.py @@ -0,0 +1,201 @@ +"""MiniMax H3 text/vision conditioning: Qwen3-VL-32B (truncated to 50 layers). + +The H3 presentation is NOT chat-templated: token ids are raw prompt/label text +(no special tokens) with explicit vision blocks spliced in: + + t2va: + fl2va: ": " [": " ] + ref2va: per condition in request order (1-based ordinals per type): + image -> ": " + audio -> "