Wayne-King/echo-memory-diffusers
0
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 