Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
omnigen.py804 linesDownload Raw Back to models
1# The code is revised from DiT2import os3import torch4import torch.nn as nn5import numpy as np6import math7from safetensors.torch import load_file8from typing import List, Optional, Tuple, Union9import torch.utils.checkpoint10from huggingface_hub import snapshot_download11from transformers.modeling_outputs import BaseModelOutputWithPast12from transformers import Phi3Config, Phi3Model13from transformers.cache_utils import Cache, DynamicCache14from transformers.utils import logging15 16 17logger = logging.get_logger(__name__)18 19 20class Phi3Transformer(Phi3Model):21    """22    Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Phi3DecoderLayer`]23    We only modified the attention mask24    Args:25        config: Phi3Config26    """27    def prefetch_layer(self, layer_idx: int, device: torch.device):28        "Starts prefetching the next layer cache"29        with torch.cuda.stream(self.prefetch_stream):30            # Prefetch next layer tensors to GPU31            for name, param in self.layers[layer_idx].named_parameters():32                param.data = param.data.to(device, non_blocking=True)33 34    def evict_previous_layer(self, layer_idx: int):35        "Moves the previous layer cache to the CPU"36        prev_layer_idx = layer_idx - 137        for name, param in self.layers[prev_layer_idx].named_parameters():38            param.data = param.data.to("cpu", non_blocking=True)39            40    def get_offlaod_layer(self, layer_idx: int, device: torch.device):41        # init stream42        if not hasattr(self, "prefetch_stream"):43            self.prefetch_stream = torch.cuda.Stream()44 45        # delete previous layer46        torch.cuda.current_stream().synchronize()47        self.evict_previous_layer(layer_idx)48        49        # make sure the current layer is ready50        torch.cuda.synchronize(self.prefetch_stream)51 52        # load next layer53        self.prefetch_layer((layer_idx + 1) % len(self.layers), device)54        55 56    def forward(57        self,58        input_ids: torch.LongTensor = None,59        attention_mask: Optional[torch.Tensor] = None,60        position_ids: Optional[torch.LongTensor] = None,61        past_key_values: Optional[List[torch.FloatTensor]] = None,62        inputs_embeds: Optional[torch.FloatTensor] = None,63        use_cache: Optional[bool] = None,64        output_attentions: Optional[bool] = None,65        output_hidden_states: Optional[bool] = None,66        return_dict: Optional[bool] = None,67        cache_position: Optional[torch.LongTensor] = None,68        offload_model: Optional[bool] = False,69    ) -> Union[Tuple, BaseModelOutputWithPast]:70        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions71        output_hidden_states = (72            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states73        )74        use_cache = use_cache if use_cache is not None else self.config.use_cache75 76        return_dict = return_dict if return_dict is not None else self.config.use_return_dict77 78        if (input_ids is None) ^ (inputs_embeds is not None):79            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")80 81        if self.gradient_checkpointing and self.training:82            if use_cache:83                logger.warning_once(84                    "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."85                )86                use_cache = False87 88        # kept for BC (non `Cache` `past_key_values` inputs)89        return_legacy_cache = False90        if use_cache and not isinstance(past_key_values, Cache):91            return_legacy_cache = True92            if past_key_values is None:93                past_key_values = DynamicCache()94            else:95                past_key_values = DynamicCache.from_legacy_cache(past_key_values)96                logger.warning_once(97                    "We detected that you are passing `past_key_values` as a tuple of tuples. This is deprecated and "98                    "will be removed in v4.47. Please convert your cache or use an appropriate `Cache` class "99                    "(https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)"100                )101 102        # if inputs_embeds is None:103        #     inputs_embeds = self.embed_tokens(input_ids)104 105        # if cache_position is None:106        #     past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0107        #     cache_position = torch.arange(108        #         past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device109        #     )110        # if position_ids is None:111        #     position_ids = cache_position.unsqueeze(0)112 113        if attention_mask is not None and attention_mask.dim() == 3:114            dtype = inputs_embeds.dtype115            min_dtype = torch.finfo(dtype).min116            attention_mask = (1 - attention_mask) * min_dtype117            attention_mask = attention_mask.unsqueeze(1).to(inputs_embeds.dtype)118        else:119            raise Exception("attention_mask parameter was unavailable or invalid")120            # causal_mask = self._update_causal_mask(121            #     attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions122            # )123 124        hidden_states = inputs_embeds125 126        # decoder layers127        all_hidden_states = () if output_hidden_states else None128        all_self_attns = () if output_attentions else None129        next_decoder_cache = None130 131        layer_idx = -1132        for decoder_layer in self.layers:133            layer_idx += 1134 135            if output_hidden_states:136                all_hidden_states += (hidden_states,)137 138            if self.gradient_checkpointing and self.training:139                layer_outputs = self._gradient_checkpointing_func(140                    decoder_layer.__call__,141                    hidden_states,142                    attention_mask,143                    position_ids,144                    past_key_values,145                    output_attentions,146                    use_cache,147                    cache_position,148                )149            else:150                if offload_model and not self.training:151                    self.get_offlaod_layer(layer_idx, device=inputs_embeds.device)152                layer_outputs = decoder_layer(153                    hidden_states,154                    attention_mask=attention_mask,155                    position_ids=position_ids,156                    past_key_value=past_key_values,157                    output_attentions=output_attentions,158                    use_cache=use_cache,159                    cache_position=cache_position,160                )161 162            hidden_states = layer_outputs[0]163 164            if use_cache:165                next_decoder_cache = layer_outputs[2 if output_attentions else 1]166 167            if output_attentions:168                all_self_attns += (layer_outputs[1],)169 170        hidden_states = self.norm(hidden_states)171 172        # add hidden states from the last decoder layer173        if output_hidden_states:174            print('************')175            all_hidden_states += (hidden_states,)176 177        next_cache = next_decoder_cache if use_cache else None178        if return_legacy_cache:179            next_cache = next_cache.to_legacy_cache()180 181        if not return_dict:182            return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)183        return BaseModelOutputWithPast(184            last_hidden_state=hidden_states,185            past_key_values=next_cache,186            hidden_states=all_hidden_states,187            attentions=all_self_attns,188        )189 190 191def modulate(x, shift, scale):192    return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)193 194 195class TimestepEmbedder(nn.Module):196    """197    Embeds scalar timesteps into vector representations.198    """199    def __init__(self, hidden_size, frequency_embedding_size=256):200        super().__init__()201        self.mlp = nn.Sequential(202            nn.Linear(frequency_embedding_size, hidden_size, bias=True),203            nn.SiLU(),204            nn.Linear(hidden_size, hidden_size, bias=True),205        )206        self.frequency_embedding_size = frequency_embedding_size207 208    @staticmethod209    def timestep_embedding(t, dim, max_period=10000):210        """211        Create sinusoidal timestep embeddings.212        :param t: a 1-D Tensor of N indices, one per batch element.213                          These may be fractional.214        :param dim: the dimension of the output.215        :param max_period: controls the minimum frequency of the embeddings.216        :return: an (N, D) Tensor of positional embeddings.217        """218        # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py219        half = dim // 2220        freqs = torch.exp(221            -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half222        ).to(device=t.device)223        args = t[:, None].float() * freqs[None]224        embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)225        if dim % 2:226            embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)227        return embedding228 229    def forward(self, t, dtype=torch.float32):230        t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(dtype)231        t_emb = self.mlp(t_freq)232        return t_emb233 234 235class FinalLayer(nn.Module):236    """237    The final layer of DiT.238    """239    def __init__(self, hidden_size, patch_size, out_channels):240        super().__init__()241        self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)242        self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)243        self.adaLN_modulation = nn.Sequential(244            nn.SiLU(),245            nn.Linear(hidden_size, 2 * hidden_size, bias=True)246        )247 248    def forward(self, x, c):249        shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)250        x = modulate(self.norm_final(x), shift, scale)251        x = self.linear(x)252        return x253 254 255def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, interpolation_scale=1.0, base_size=1):256    """257    grid_size: int of the grid height and width return: pos_embed: [grid_size*grid_size, embed_dim] or258    [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)259    """260    if isinstance(grid_size, int):261        grid_size = (grid_size, grid_size)262 263    grid_h = np.arange(grid_size[0], dtype=np.float32) / (grid_size[0] / base_size) / interpolation_scale264    grid_w = np.arange(grid_size[1], dtype=np.float32) / (grid_size[1] / base_size) / interpolation_scale265    grid = np.meshgrid(grid_w, grid_h)  # here w goes first266    grid = np.stack(grid, axis=0)267 268    grid = grid.reshape([2, 1, grid_size[1], grid_size[0]])269    pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)270    if cls_token and extra_tokens > 0:271        pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)272    return pos_embed273 274 275def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):276    assert embed_dim % 2 == 0277 278    # use half of dimensions to encode grid_h279    emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0])  # (H*W, D/2)280    emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1])  # (H*W, D/2)281 282    emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)283    return emb284 285 286def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):287    """288    embed_dim: output dimension for each position289    pos: a list of positions to be encoded: size (M,)290    out: (M, D)291    """292    assert embed_dim % 2 == 0293    omega = np.arange(embed_dim // 2, dtype=np.float64)294    omega /= embed_dim / 2.295    omega = 1. / 10000**omega  # (D/2,)296 297    pos = pos.reshape(-1)  # (M,)298    out = np.einsum('m,d->md', pos, omega)  # (M, D/2), outer product299 300    emb_sin = np.sin(out) # (M, D/2)301    emb_cos = np.cos(out) # (M, D/2)302 303    emb = np.concatenate([emb_sin, emb_cos], axis=1)  # (M, D)304    return emb305 306 307class PatchEmbedMR(nn.Module):308    """ 2D Image to Patch Embedding309    """310    def __init__(311            self,312            patch_size: int = 2,313            in_chans: int = 4,314            embed_dim: int = 768,315            bias: bool = True,316    ):317        super().__init__()318        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)319 320    def forward(self, x):321        x = self.proj(x)322        x = x.flatten(2).transpose(1, 2)  # NCHW -> NLC323        return x324 325 326class OmniGenOriginalModel(nn.Module):327    """328    Diffusion model with a Transformer backbone.329    """330    def __init__(331        self,332        transformer_config: Phi3Config,333        patch_size=2,334        in_channels=4,335        pe_interpolation: float = 1.0,336        pos_embed_max_size: int = 192,337    ):338        super().__init__()339        self.in_channels = in_channels340        self.out_channels = in_channels341        self.patch_size = patch_size342        self.pos_embed_max_size = pos_embed_max_size343 344        hidden_size = transformer_config.hidden_size345 346        self.x_embedder = PatchEmbedMR(patch_size, in_channels, hidden_size, bias=True)347        self.input_x_embedder = PatchEmbedMR(patch_size, in_channels, hidden_size, bias=True)348 349        self.time_token = TimestepEmbedder(hidden_size)350        self.t_embedder = TimestepEmbedder(hidden_size)351        352        self.pe_interpolation = pe_interpolation353        pos_embed = get_2d_sincos_pos_embed(hidden_size, pos_embed_max_size, interpolation_scale=self.pe_interpolation, base_size=64)354        self.register_buffer("pos_embed", torch.from_numpy(pos_embed).float().unsqueeze(0), persistent=True)355 356        self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)357 358        self.initialize_weights()359 360        self.llm = Phi3Transformer(config=transformer_config)361        self.llm.config.use_cache = False362    363    @classmethod364    def from_pretrained(cls, model_name):365        if not os.path.exists(model_name):366            cache_folder = os.getenv('HF_HUB_CACHE')367            model_name = snapshot_download(repo_id=model_name,368                                           cache_dir=cache_folder,369                                           ignore_patterns=['flax_model.msgpack', 'rust_model.ot', 'tf_model.h5'])370        config = Phi3Config.from_pretrained(model_name)371        model = cls(config)372        if os.path.exists(os.path.join(model_name, 'model.safetensors')):373            print("Loading safetensors")374            ckpt = load_file(os.path.join(model_name, 'model.safetensors'))375        else:376            ckpt = torch.load(os.path.join(model_name, 'model.pt'), map_location='cpu')377        model.load_state_dict(ckpt)378        return model379 380    def initialize_weights(self):381        assert not hasattr(self, "llama")382 383        # Initialize transformer layers:384        def _basic_init(module):385            if isinstance(module, nn.Linear):386                torch.nn.init.xavier_uniform_(module.weight)387                if module.bias is not None:388                    nn.init.constant_(module.bias, 0)389        self.apply(_basic_init)390        391        # Initialize patch_embed like nn.Linear (instead of nn.Conv2d):392        w = self.x_embedder.proj.weight.data393        nn.init.xavier_uniform_(w.view([w.shape[0], -1]))394        nn.init.constant_(self.x_embedder.proj.bias, 0)395 396        w = self.input_x_embedder.proj.weight.data397        nn.init.xavier_uniform_(w.view([w.shape[0], -1]))398        nn.init.constant_(self.x_embedder.proj.bias, 0)399 400 401        # Initialize timestep embedding MLP:402        nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)403        nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)404        nn.init.normal_(self.time_token.mlp[0].weight, std=0.02)405        nn.init.normal_(self.time_token.mlp[2].weight, std=0.02)406 407        # Zero-out output layers:408        nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)409        nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)410        nn.init.constant_(self.final_layer.linear.weight, 0)411        nn.init.constant_(self.final_layer.linear.bias, 0)412 413    def unpatchify(self, x, h, w):414        """415        x: (N, T, patch_size**2 * C)416        imgs: (N, H, W, C)417        """418        c = self.out_channels419 420        x = x.reshape(shape=(x.shape[0], h//self.patch_size, w//self.patch_size, self.patch_size, self.patch_size, c))421        x = torch.einsum('nhwpqc->nchpwq', x)422        imgs = x.reshape(shape=(x.shape[0], c, h, w))423        return imgs424 425 426    def cropped_pos_embed(self, height, width):427        """Crops positional embeddings for SD3 compatibility."""428        if self.pos_embed_max_size is None:429            raise ValueError("`pos_embed_max_size` must be set for cropping.")430 431        height = height // self.patch_size432        width = width // self.patch_size433        if height > self.pos_embed_max_size:434            raise ValueError(435                f"Height ({height}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}."436            )437        if width > self.pos_embed_max_size:438            raise ValueError(439                f"Width ({width}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}."440            )441 442        top = (self.pos_embed_max_size - height) // 2443        left = (self.pos_embed_max_size - width) // 2444        spatial_pos_embed = self.pos_embed.reshape(1, self.pos_embed_max_size, self.pos_embed_max_size, -1)445        spatial_pos_embed = spatial_pos_embed[:, top : top + height, left : left + width, :]446        # print(top, top + height, left, left + width, spatial_pos_embed.size())447        spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1])448        return spatial_pos_embed449 450 451    def patch_multiple_resolutions(self, latents, padding_latent=None, is_input_images:bool=False):452        if isinstance(latents, list):453            return_list = False454            if padding_latent is None:455                padding_latent = [None] * len(latents)456                return_list = True457            patched_latents, num_tokens, shapes = [], [], []458            for latent, padding in zip(latents, padding_latent):459                height, width = latent.shape[-2:]460                if is_input_images:461                    latent = self.input_x_embedder(latent)462                else:463                    latent = self.x_embedder(latent)464                pos_embed = self.cropped_pos_embed(height, width)    465                latent = latent + pos_embed466                if padding is not None:467                    latent = torch.cat([latent, padding], dim=-2)468                patched_latents.append(latent)469 470                num_tokens.append(pos_embed.size(1))471                shapes.append([height, width])472            if not return_list:473                latents = torch.cat(patched_latents, dim=0)474            else:475                latents = patched_latents476        else:477            height, width = latents.shape[-2:]478            if is_input_images:479                latents = self.input_x_embedder(latents)480            else:481                latents = self.x_embedder(latents)482            pos_embed = self.cropped_pos_embed(height, width)  483            latents = latents + pos_embed484            num_tokens = latents.size(1)485            shapes = [height, width]486        return latents, num_tokens, shapes487 488    489    def forward(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, padding_latent=None, past_key_values=None, return_past_key_values=True, offload_model:bool=False):490        """491        492        """493        input_is_list = isinstance(x, list)494        x, num_tokens, shapes = self.patch_multiple_resolutions(x, padding_latent)495        time_token = self.time_token(timestep, dtype=x[0].dtype).unsqueeze(1)   496        497        if input_img_latents is not None:498            input_latents, _, _ = self.patch_multiple_resolutions(input_img_latents, is_input_images=True)499        if input_ids is not None:500            condition_embeds = self.llm.embed_tokens(input_ids).clone()501            input_img_inx = 0502            for b_inx in input_image_sizes.keys():503                for start_inx, end_inx in input_image_sizes[b_inx]:504                    condition_embeds[b_inx, start_inx: end_inx] = input_latents[input_img_inx]505                    input_img_inx += 1506            if input_img_latents is not None:507                assert input_img_inx == len(input_latents) 508 509            input_emb = torch.cat([condition_embeds, time_token, x], dim=1)510        else:511            input_emb = torch.cat([time_token, x], dim=1)512        output = self.llm(inputs_embeds=input_emb, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, offload_model=offload_model)513        output, past_key_values = output.last_hidden_state, output.past_key_values514        if input_is_list:515            image_embedding = output[:, -max(num_tokens):]516            time_emb = self.t_embedder(timestep, dtype=x.dtype)517            x = self.final_layer(image_embedding, time_emb)518            latents = []519            for i in range(x.size(0)):520                latent = x[i:i+1, :num_tokens[i]]521                latent = self.unpatchify(latent, shapes[i][0], shapes[i][1])522                latents.append(latent)523        else:524            image_embedding = output[:, -num_tokens:]525            time_emb = self.t_embedder(timestep, dtype=x.dtype)526            x = self.final_layer(image_embedding, time_emb)527            latents = self.unpatchify(x, shapes[0], shapes[1])528 529        if return_past_key_values:530            return latents, past_key_values531        return latents532 533    @torch.no_grad()534    def forward_with_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):      535        self.llm.config.use_cache = use_kv_cache536        model_out, past_key_values = self.forward(x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, past_key_values=past_key_values, return_past_key_values=True, offload_model=offload_model)537        if use_img_cfg:538            cond, uncond, img_cond = torch.split(model_out, len(model_out) // 3, dim=0)539            cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)540            model_out = [cond, cond, cond]541        else:542            cond, uncond = torch.split(model_out, len(model_out) // 2, dim=0)543            cond = uncond + cfg_scale * (cond - uncond)544            model_out = [cond, cond]545        546        return torch.cat(model_out, dim=0), past_key_values547 548 549    @torch.no_grad()550    def forward_with_separate_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):551        self.llm.config.use_cache = use_kv_cache552        if past_key_values is None:553            past_key_values = [None] * len(attention_mask)554 555        x = torch.split(x, len(x) // len(attention_mask), dim=0)556        timestep = timestep.to(x[0].dtype)557        timestep = torch.split(timestep, len(timestep) // len(input_ids), dim=0)558 559        model_out, pask_key_values = [], []560        for i in range(len(input_ids)):561            temp_out, temp_pask_key_values = self.forward(x[i], timestep[i], input_ids[i], input_img_latents[i], input_image_sizes[i], attention_mask[i], position_ids[i], past_key_values=past_key_values[i], return_past_key_values=True, offload_model=offload_model)562            model_out.append(temp_out)563            pask_key_values.append(temp_pask_key_values)564 565        if len(model_out) == 3:566            cond, uncond, img_cond = model_out567            cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)568            model_out = [cond, cond, cond]569        elif len(model_out) == 2:570            cond, uncond = model_out571            cond = uncond + cfg_scale * (cond - uncond)572            model_out = [cond, cond]573        else:574            return model_out[0]575        576        return torch.cat(model_out, dim=0), pask_key_values577 578 579 580class OmniGenTransformer(OmniGenOriginalModel):581    def __init__(self):582        config = {583            "_name_or_path": "Phi-3-vision-128k-instruct",584            "architectures": [585                "Phi3ForCausalLM"586            ],587            "attention_dropout": 0.0,588            "bos_token_id": 1,589            "eos_token_id": 2,590            "hidden_act": "silu",591            "hidden_size": 3072,592            "initializer_range": 0.02,593            "intermediate_size": 8192,594            "max_position_embeddings": 131072,595            "model_type": "phi3",596            "num_attention_heads": 32,597            "num_hidden_layers": 32,598            "num_key_value_heads": 32,599            "original_max_position_embeddings": 4096,600            "rms_norm_eps": 1e-05,601            "rope_scaling": {602                "long_factor": [603                1.0299999713897705,604                1.0499999523162842,605                1.0499999523162842,606                1.0799999237060547,607                1.2299998998641968,608                1.2299998998641968,609                1.2999999523162842,610                1.4499999284744263,611                1.5999999046325684,612                1.6499998569488525,613                1.8999998569488525,614                2.859999895095825,615                3.68999981880188,616                5.419999599456787,617                5.489999771118164,618                5.489999771118164,619                9.09000015258789,620                11.579999923706055,621                15.65999984741211,622                15.769999504089355,623                15.789999961853027,624                18.360000610351562,625                21.989999771118164,626                23.079999923706055,627                30.009998321533203,628                32.35000228881836,629                32.590003967285156,630                35.56000518798828,631                39.95000457763672,632                53.840003967285156,633                56.20000457763672,634                57.95000457763672,635                59.29000473022461,636                59.77000427246094,637                59.920005798339844,638                61.190006256103516,639                61.96000671386719,640                62.50000762939453,641                63.3700065612793,642                63.48000717163086,643                63.48000717163086,644                63.66000747680664,645                63.850006103515625,646                64.08000946044922,647                64.760009765625,648                64.80001068115234,649                64.81001281738281,650                64.81001281738281651                ],652                "short_factor": [653                1.05,654                1.05,655                1.05,656                1.1,657                1.1,658                1.1,659                1.2500000000000002,660                1.2500000000000002,661                1.4000000000000004,662                1.4500000000000004,663                1.5500000000000005,664                1.8500000000000008,665                1.9000000000000008,666                2.000000000000001,667                2.000000000000001,668                2.000000000000001,669                2.000000000000001,670                2.000000000000001,671                2.000000000000001,672                2.000000000000001,673                2.000000000000001,674                2.000000000000001,675                2.000000000000001,676                2.000000000000001,677                2.000000000000001,678                2.000000000000001,679                2.000000000000001,680                2.000000000000001,681                2.000000000000001,682                2.000000000000001,683                2.000000000000001,684                2.000000000000001,685                2.1000000000000005,686                2.1000000000000005,687                2.2,688                2.3499999999999996,689                2.3499999999999996,690                2.3499999999999996,691                2.3499999999999996,692                2.3999999999999995,693                2.3999999999999995,694                2.6499999999999986,695                2.6999999999999984,696                2.8999999999999977,697                2.9499999999999975,698                3.049999999999997,699                3.049999999999997,700                3.049999999999997701                ],702                "type": "su"703            },704            "rope_theta": 10000.0,705            "sliding_window": 131072,706            "tie_word_embeddings": False,707            "torch_dtype": "bfloat16",708            "transformers_version": "4.38.1",709            "use_cache": True,710            "vocab_size": 32064,711            "_attn_implementation": "sdpa"712        }713        config = Phi3Config(**config)714        super().__init__(config)715 716    717    def forward(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, padding_latent=None, past_key_values=None, return_past_key_values=True, offload_model:bool=False):718        input_is_list = isinstance(x, list)719        x, num_tokens, shapes = self.patch_multiple_resolutions(x, padding_latent)720        time_token = self.time_token(timestep, dtype=x[0].dtype).unsqueeze(1)   721        722        if input_img_latents is not None:723            input_latents, _, _ = self.patch_multiple_resolutions(input_img_latents, is_input_images=True)724        if input_ids is not None:725            condition_embeds = self.llm.embed_tokens(input_ids).clone()726            input_img_inx = 0727            for b_inx in input_image_sizes.keys():728                for start_inx, end_inx in input_image_sizes[b_inx]:729                    condition_embeds[b_inx, start_inx: end_inx] = input_latents[input_img_inx]730                    input_img_inx += 1731            if input_img_latents is not None:732                assert input_img_inx == len(input_latents) 733 734            input_emb = torch.cat([condition_embeds, time_token, x], dim=1)735        else:736            input_emb = torch.cat([time_token, x], dim=1)737        output = self.llm(inputs_embeds=input_emb, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, offload_model=offload_model)738        output, past_key_values = output.last_hidden_state, output.past_key_values739        if input_is_list:740            image_embedding = output[:, -max(num_tokens):]741            time_emb = self.t_embedder(timestep, dtype=x.dtype)742            x = self.final_layer(image_embedding, time_emb)743            latents = []744            for i in range(x.size(0)):745                latent = x[i:i+1, :num_tokens[i]]746                latent = self.unpatchify(latent, shapes[i][0], shapes[i][1])747                latents.append(latent)748        else:749            image_embedding = output[:, -num_tokens:]750            time_emb = self.t_embedder(timestep, dtype=x.dtype)751            x = self.final_layer(image_embedding, time_emb)752            latents = self.unpatchify(x, shapes[0], shapes[1])753 754        if return_past_key_values:755            return latents, past_key_values756        return latents757    758 759    @torch.no_grad()760    def forward_with_separate_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):761        self.llm.config.use_cache = use_kv_cache762        if past_key_values is None:763            past_key_values = [None] * len(attention_mask)764 765        x = torch.split(x, len(x) // len(attention_mask), dim=0)766        timestep = timestep.to(x[0].dtype)767        timestep = torch.split(timestep, len(timestep) // len(input_ids), dim=0)768 769        model_out, pask_key_values = [], []770        for i in range(len(input_ids)):771            temp_out, temp_pask_key_values = self.forward(x[i], timestep[i], input_ids[i], input_img_latents[i], input_image_sizes[i], attention_mask[i], position_ids[i], past_key_values=past_key_values[i], return_past_key_values=True, offload_model=offload_model)772            model_out.append(temp_out)773            pask_key_values.append(temp_pask_key_values)774 775        if len(model_out) == 3:776            cond, uncond, img_cond = model_out777            cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)778            model_out = [cond, cond, cond]779        elif len(model_out) == 2:780            cond, uncond = model_out781            cond = uncond + cfg_scale * (cond - uncond)782            model_out = [cond, cond]783        else:784            return model_out[0]785        786        return torch.cat(model_out, dim=0), pask_key_values787    788 789    @staticmethod790    def state_dict_converter():791        return OmniGenTransformerStateDictConverter()792 793 794 795class OmniGenTransformerStateDictConverter:796    def __init__(self):797        pass798 799    def from_diffusers(self, state_dict):800        return state_dict801    802    def from_civitai(self, state_dict):803        return state_dict804