hugging-apps/echo-memory
0
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 