hugging-apps/echo-memory
0
1from transformers import BertModel, BertConfig, T5EncoderModel, T5Config2import torch3 4 5 6class HunyuanDiTCLIPTextEncoder(BertModel):7 def __init__(self):8 config = BertConfig(9 _name_or_path = "",10 architectures = ["BertModel"],11 attention_probs_dropout_prob = 0.1,12 bos_token_id = 0,13 classifier_dropout = None,14 directionality = "bidi",15 eos_token_id = 2,16 hidden_act = "gelu",17 hidden_dropout_prob = 0.1,18 hidden_size = 1024,19 initializer_range = 0.02,20 intermediate_size = 4096,21 layer_norm_eps = 1e-12,22 max_position_embeddings = 512,23 model_type = "bert",24 num_attention_heads = 16,25 num_hidden_layers = 24,26 output_past = True,27 pad_token_id = 0,28 pooler_fc_size = 768,29 pooler_num_attention_heads = 12,30 pooler_num_fc_layers = 3,31 pooler_size_per_head = 128,32 pooler_type = "first_token_transform",33 position_embedding_type = "absolute",34 torch_dtype = "float32",35 transformers_version = "4.37.2",36 type_vocab_size = 2,37 use_cache = True,38 vocab_size = 4702039 )40 super().__init__(config, add_pooling_layer=False)41 self.eval()42 43 def forward(self, input_ids, attention_mask, clip_skip=1):44 input_shape = input_ids.size()45 46 batch_size, seq_length = input_shape47 device = input_ids.device48 49 past_key_values_length = 050 51 if attention_mask is None:52 attention_mask = torch.ones(((batch_size, seq_length + past_key_values_length)), device=device)53 54 extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(attention_mask, input_shape)55 56 embedding_output = self.embeddings(57 input_ids=input_ids,58 position_ids=None,59 token_type_ids=None,60 inputs_embeds=None,61 past_key_values_length=0,62 )63 encoder_outputs = self.encoder(64 embedding_output,65 attention_mask=extended_attention_mask,66 head_mask=None,67 encoder_hidden_states=None,68 encoder_attention_mask=None,69 past_key_values=None,70 use_cache=False,71 output_attentions=False,72 output_hidden_states=True,73 return_dict=True,74 )75 all_hidden_states = encoder_outputs.hidden_states76 prompt_emb = all_hidden_states[-clip_skip]77 if clip_skip > 1:78 mean, std = all_hidden_states[-1].mean(), all_hidden_states[-1].std()79 prompt_emb = (prompt_emb - prompt_emb.mean()) / prompt_emb.std() * std + mean80 return prompt_emb81 82 @staticmethod83 def state_dict_converter():84 return HunyuanDiTCLIPTextEncoderStateDictConverter()85 86 87 88class HunyuanDiTT5TextEncoder(T5EncoderModel):89 def __init__(self):90 config = T5Config(91 _name_or_path = "../HunyuanDiT/t2i/mt5",92 architectures = ["MT5ForConditionalGeneration"],93 classifier_dropout = 0.0,94 d_ff = 5120,95 d_kv = 64,96 d_model = 2048,97 decoder_start_token_id = 0,98 dense_act_fn = "gelu_new",99 dropout_rate = 0.1,100 eos_token_id = 1,101 feed_forward_proj = "gated-gelu",102 initializer_factor = 1.0,103 is_encoder_decoder = True,104 is_gated_act = True,105 layer_norm_epsilon = 1e-06,106 model_type = "t5",107 num_decoder_layers = 24,108 num_heads = 32,109 num_layers = 24,110 output_past = True,111 pad_token_id = 0,112 relative_attention_max_distance = 128,113 relative_attention_num_buckets = 32,114 tie_word_embeddings = False,115 tokenizer_class = "T5Tokenizer",116 transformers_version = "4.37.2",117 use_cache = True,118 vocab_size = 250112119 )120 super().__init__(config)121 self.eval()122 123 def forward(self, input_ids, attention_mask, clip_skip=1):124 outputs = super().forward(125 input_ids=input_ids,126 attention_mask=attention_mask,127 output_hidden_states=True,128 )129 prompt_emb = outputs.hidden_states[-clip_skip]130 if clip_skip > 1:131 mean, std = outputs.hidden_states[-1].mean(), outputs.hidden_states[-1].std()132 prompt_emb = (prompt_emb - prompt_emb.mean()) / prompt_emb.std() * std + mean133 return prompt_emb134 135 @staticmethod136 def state_dict_converter():137 return HunyuanDiTT5TextEncoderStateDictConverter()138 139 140 141class HunyuanDiTCLIPTextEncoderStateDictConverter():142 def __init__(self):143 pass144 145 def from_diffusers(self, state_dict):146 state_dict_ = {name[5:]: param for name, param in state_dict.items() if name.startswith("bert.")}147 return state_dict_148 149 def from_civitai(self, state_dict):150 return self.from_diffusers(state_dict)151 152 153class HunyuanDiTT5TextEncoderStateDictConverter():154 def __init__(self):155 pass156 157 def from_diffusers(self, state_dict):158 state_dict_ = {name: param for name, param in state_dict.items() if name.startswith("encoder.")}159 state_dict_["shared.weight"] = state_dict["shared.weight"]160 return state_dict_161 162 def from_civitai(self, state_dict):163 return self.from_diffusers(state_dict)164 