hugging-apps/echo-memory
0
1import torch2from diffsynth.models.svd_unet import TemporalTimesteps3 4 5class MultiValueEncoder(torch.nn.Module):6 def __init__(self, encoders=()):7 super().__init__()8 self.encoders = torch.nn.ModuleList(encoders)9 10 def __call__(self, values, dtype):11 emb = []12 for encoder, value in zip(self.encoders, values):13 if value is not None:14 value = value.unsqueeze(0)15 emb.append(encoder(value, dtype))16 emb = torch.concat(emb, dim=0)17 return emb18 19 20class SingleValueEncoder(torch.nn.Module):21 def __init__(self, dim_in=256, dim_out=4096, prefer_len=32, computation_device=None):22 super().__init__()23 self.prefer_len = prefer_len24 self.prefer_proj = TemporalTimesteps(num_channels=dim_in, flip_sin_to_cos=True, downscale_freq_shift=0, computation_device=computation_device)25 self.prefer_value_embedder = torch.nn.Sequential(26 torch.nn.Linear(dim_in, dim_out), torch.nn.SiLU(), torch.nn.Linear(dim_out, dim_out)27 )28 self.positional_embedding = torch.nn.Parameter(29 torch.randn(self.prefer_len, dim_out) 30 )31 self._initialize_weights()32 33 def _initialize_weights(self):34 last_linear = self.prefer_value_embedder[-1]35 torch.nn.init.zeros_(last_linear.weight)36 torch.nn.init.zeros_(last_linear.bias)37 38 def forward(self, value, dtype):39 value = value * 100040 emb = self.prefer_proj(value).to(dtype)41 emb = self.prefer_value_embedder(emb).squeeze(0)42 base_embeddings = emb.expand(self.prefer_len, -1)43 positional_embedding = self.positional_embedding.to(dtype=base_embeddings.dtype, device=base_embeddings.device)44 learned_embeddings = base_embeddings + positional_embedding45 return learned_embeddings46 47 @staticmethod48 def state_dict_converter():49 return SingleValueEncoderStateDictConverter()50 51 52class SingleValueEncoderStateDictConverter:53 def __init__(self):54 pass55 56 def from_diffusers(self, state_dict):57 return state_dict58 59 def from_civitai(self, state_dict):60 return state_dict61 