Instructions to use drozbay/Wan2.2-S2V-14B-module with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Wan2.2
How to use drozbay/Wan2.2-S2V-14B-module with Wan2.2:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
| """Bernini in-context conditioning (context_latents) support for Wan S2V models. | |
| Adds a "Wan S2V-Bernini Patch" node (MODEL -> MODEL). Insert it after the model | |
| loader when running a Bernini trunk with the Wan2.2-S2V module stacked on (e.g. | |
| via DiffusionModelLoaderKJ with extra_state_dict), and set conditioning with the | |
| native BerniniConditioning node (optionally combined with WanSoundImageToVideo | |
| for audio). | |
| Uses only official ComfyUI extension points, no monkey patching: | |
| - WanModel._forward (inherited by S2V) already pads context_latents and appends | |
| their rope freqs; the native S2V forward_orig just never appends the tokens, | |
| which crashes attention on the freqs length mismatch. | |
| - A DIFFUSION_MODEL wrapper patch-embeds the context latents (it can reach the | |
| model via executor.class_obj) and stashes the tokens in transformer_options. | |
| - A dit block-0 replace patch inserts the tokens right after the main video | |
| tokens, before the first attention runs, restoring the token/freqs alignment. | |
| They ride through all blocks (audio injection writes only x[:, :seq_len]) and | |
| unpatchify() drops them at the output. | |
| When context_latents cannot be honored (non-S2V model, wrong latent channels), | |
| they are removed from the call with a warning instead of crashing attention. | |
| Models with native context_latents support (base WanModel, or a future S2V | |
| implementation) are passed through untouched. | |
| """ | |
| import inspect | |
| import logging | |
| import torch | |
| import comfy.ldm.common_dit | |
| from comfy.ldm.wan import model as wan | |
| from comfy.patcher_extension import WrappersMP | |
| _TOKENS = "_wan_s2v_bernini_tokens" | |
| _MAIN_LEN = "_wan_s2v_bernini_main_len" | |
| _native_support = {} | |
| def _has_native_support(cls): | |
| if cls not in _native_support: | |
| try: | |
| _native_support[cls] = "context_latents" in inspect.getsource(cls.forward_orig) | |
| except (OSError, TypeError): | |
| _native_support[cls] = False | |
| return _native_support[cls] | |
| def _drop_reason(model, context_latents): | |
| cls = type(model) | |
| if cls is not wan.WanModel_S2V: | |
| return f"{cls.__name__} is not supported" | |
| if cls._forward is not wan.WanModel._forward: | |
| return "WanModel_S2V no longer inherits WanModel._forward" | |
| ch = model.patch_embedding.weight.shape[1] | |
| if any(lat.shape[1] != ch for lat in context_latents): | |
| return f"model expects {ch}-channel latents" | |
| return None | |
| def _embed_context_latents(executor, x, timestep, context, clip_fea=None, time_dim_concat=None, transformer_options={}, **kwargs): | |
| model = executor.class_obj | |
| context_latents = kwargs.get("context_latents", None) | |
| transformer_options.pop(_TOKENS, None) | |
| if context_latents is not None and not _has_native_support(type(model)): | |
| reason = _drop_reason(model, context_latents) | |
| if reason is not None: | |
| logging.warning(f"wan_s2v_module_patch: dropping context_latents: {reason}") | |
| kwargs = {k: v for k, v in kwargs.items() if k != "context_latents"} | |
| else: | |
| p = model.patch_size | |
| xp = comfy.ldm.common_dit.pad_to_patch_size(x, p) | |
| t_len = xp.shape[-3] | |
| if time_dim_concat is not None: | |
| t_len += comfy.ldm.common_dit.pad_to_patch_size(time_dim_concat, p).shape[-3] | |
| transformer_options[_MAIN_LEN] = (t_len // p[0]) * (xp.shape[-2] // p[1]) * (xp.shape[-1] // p[2]) | |
| tokens = [] | |
| for lat in context_latents: | |
| lat = comfy.ldm.common_dit.pad_to_patch_size(lat, p) | |
| tokens.append(model.patch_embedding(lat.float().to(x.device)).flatten(2).transpose(1, 2)) | |
| transformer_options[_TOKENS] = torch.cat(tokens, dim=1).to(x.dtype) | |
| return executor(x, timestep, context, clip_fea, time_dim_concat, transformer_options, **kwargs) | |
| def _insert_tokens_block0(args, extra): | |
| transformer_options = args["transformer_options"] | |
| tokens = transformer_options.get(_TOKENS, None) | |
| if tokens is not None: | |
| n = transformer_options[_MAIN_LEN] | |
| img = args["img"] | |
| args = {**args, "img": torch.cat([img[:, :n], tokens.to(img.dtype), img[:, n:]], dim=1)} | |
| return extra["original_block"](args) | |
| class WanS2VBerniniPatch: | |
| def INPUT_TYPES(cls): | |
| return {"required": {"model": ("MODEL",)}} | |
| RETURN_TYPES = ("MODEL",) | |
| FUNCTION = "patch" | |
| CATEGORY = "model_patches/video" | |
| DESCRIPTION = "Enables Bernini context_latents conditioning (BerniniConditioning node) on Wan S2V models, e.g. a Bernini trunk with the S2V module stacked on. No effect on other models or without BerniniConditioning." | |
| def patch(self, model): | |
| m = model.clone() | |
| m.add_wrapper_with_key(WrappersMP.DIFFUSION_MODEL, "wan_s2v_bernini", _embed_context_latents) | |
| m.set_model_patch_replace(_insert_tokens_block0, "dit", "double_block", 0) | |
| return (m,) | |
| NODE_CLASS_MAPPINGS = {"WanS2VBerniniPatch": WanS2VBerniniPatch} | |
| NODE_DISPLAY_NAME_MAPPINGS = {"WanS2VBerniniPatch": "Wan S2V-Bernini Patch"} | |