Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
hunyuan_dit_text_encoder.py164 linesDownload Raw Back to models
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