Team Ai
Modelpublic

Wayne-King/echo-memory-diffusers

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes
pipeline.py179 linesDownload Raw Back to root
1# Copyright 2026 Echo Team and The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15"""Echo-Memory community pipeline for official Wan 2.1 Diffusers weights.16 17Loads `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`, then overlays the released18`context_k1` row from `Echo-Team/Echo-Memory` after remapping original19DiffSynth / Wan keys onto the Diffusers transformer.20 21This is a community overlay, not a new official Wan checkpoint. Extra22action-MLP / SSM slots stay in the Echo-Memory research stack.23 24Paper: https://arxiv.org/abs/2606.0980325Code: https://github.com/Echo-Team-Joy-Future-Academy-JD/Echo-Memory26"""27 28from typing import Dict, Iterable, List, Optional, Tuple29 30import torch31from huggingface_hub import hf_hub_download32from safetensors.torch import load_file33 34from diffusers import WanPipeline35from diffusers.utils import logging36 37 38logger = logging.get_logger(__name__)  # pylint: disable=invalid-name39 40DEFAULT_BASE_MODEL = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"41DEFAULT_REPO_ID = "Echo-Team/Echo-Memory"42DEFAULT_FILENAME = "context_k1/epoch-0.safetensors"43DEFAULT_CONVERTED_REPO_ID = "Wayne-King/echo-memory-diffusers"44DEFAULT_CONVERTED_FILENAME = "context_k1-diffusers/diffusion_pytorch_model.safetensors"45 46SKIP_SUBSTRINGS = (47    "action_mlp",48    "self_attn_with_action",49    "block_wise_ssm",50    "videossm_hybrid",51    "spatial_memory_module",52)53 54# Same mapping as `scripts/convert_wan_to_diffusers.py` for Wan 2.1 T2V.55# Duplicated here because that script is not an importable package.56TRANSFORMER_KEYS_RENAME_DICT = {57    "time_embedding.0": "condition_embedder.time_embedder.linear_1",58    "time_embedding.2": "condition_embedder.time_embedder.linear_2",59    "text_embedding.0": "condition_embedder.text_embedder.linear_1",60    "text_embedding.2": "condition_embedder.text_embedder.linear_2",61    "time_projection.1": "condition_embedder.time_proj",62    "head.modulation": "scale_shift_table",63    "head.head": "proj_out",64    "modulation": "scale_shift_table",65    "ffn.0": "ffn.net.0.proj",66    "ffn.2": "ffn.net.2",67    # The original model names norms as norm1, norm3, norm2.68    # Diffusers uses norm1, norm2, norm3.69    "norm2": "norm__placeholder",70    "norm3": "norm2",71    "norm__placeholder": "norm3",72    "self_attn.q": "attn1.to_q",73    "self_attn.k": "attn1.to_k",74    "self_attn.v": "attn1.to_v",75    "self_attn.o": "attn1.to_out.0",76    "self_attn.norm_q": "attn1.norm_q",77    "self_attn.norm_k": "attn1.norm_k",78    "cross_attn.q": "attn2.to_q",79    "cross_attn.k": "attn2.to_k",80    "cross_attn.v": "attn2.to_v",81    "cross_attn.o": "attn2.to_out.0",82    "cross_attn.norm_q": "attn2.norm_q",83    "cross_attn.norm_k": "attn2.norm_k",84}85 86 87def is_diffusers_transformer_state_dict(keys: Iterable[str]) -> bool:88    keys = list(keys)89    return any(key.startswith("condition_embedder.") or ".attn1." in key for key in keys)90 91 92def convert_echo_memory_transformer_state_dict(93    state_dict: Dict[str, torch.Tensor],94    skip_substrings: Iterable[str] = SKIP_SUBSTRINGS,95) -> Tuple[Dict[str, torch.Tensor], List[str]]:96    """Convert original Echo-Memory / DiffSynth Wan keys to Diffusers names."""97    skip_substrings = tuple(skip_substrings)98    if is_diffusers_transformer_state_dict(state_dict):99        converted = {100            key: value101            for key, value in state_dict.items()102            if not any(token in key for token in skip_substrings)103        }104        skipped = [key for key in state_dict if key not in converted]105        return converted, skipped106 107    converted = {}108    skipped = []109    for key, value in state_dict.items():110        if any(token in key for token in skip_substrings):111            skipped.append(key)112            continue113        new_key = key114        for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items():115            new_key = new_key.replace(replace_key, rename_key)116        converted[new_key] = value117    return converted, skipped118 119 120class EchoMemoryPipeline(WanPipeline):121    """Wan 2.1 T2V pipeline with an Echo-Memory `context_k1` overlay.122 123    `load_echo_memory_weights` replaces `self.transformer` parameters in place.124    Call it once after `from_pretrained`, before generation.125    """126 127    def load_echo_memory_weights(128        self,129        repo_id: str = DEFAULT_REPO_ID,130        filename: str = DEFAULT_FILENAME,131        local_path: Optional[str] = None,132        strict: bool = False,133    ):134        """Download one Echo-Memory row and overlay it on `self.transformer`."""135        if getattr(self, "transformer", None) is None:136            raise ValueError("pipeline.transformer is empty; load Wan 2.1 1.3B before overlaying Echo-Memory.")137 138        ckpt_path = local_path or hf_hub_download(repo_id=repo_id, filename=filename)139        raw = load_file(ckpt_path)140        converted, skipped = convert_echo_memory_transformer_state_dict(raw)141        missing, unexpected = self.transformer.load_state_dict(converted, strict=strict)142        logger.info(143            "Overlaid %s/%s transformer keys from %s (skipped=%s, missing=%s, unexpected=%s)",144            len(converted),145            len(raw),146            ckpt_path,147            len(skipped),148            len(missing),149            len(unexpected),150        )151        return missing, unexpected, skipped152 153    def load_converted_echo_memory_weights(154        self,155        repo_id: str = DEFAULT_CONVERTED_REPO_ID,156        filename: str = DEFAULT_CONVERTED_FILENAME,157        local_path: Optional[str] = None,158        strict: bool = False,159    ):160        """Overlay the already-remapped `context_k1` transformer weights."""161        return self.load_echo_memory_weights(162            repo_id=repo_id,163            filename=filename,164            local_path=local_path,165            strict=strict,166        )167 168    @classmethod169    def from_echo_memory(170        cls,171        pretrained_model_name_or_path: str = DEFAULT_BASE_MODEL,172        echo_memory_repo: str = DEFAULT_REPO_ID,173        echo_memory_filename: str = DEFAULT_FILENAME,174        **kwargs,175    ):176        pipe = cls.from_pretrained(pretrained_model_name_or_path, **kwargs)177        pipe.load_echo_memory_weights(repo_id=echo_memory_repo, filename=echo_memory_filename)178        return pipe179