hugging-apps/echo-memory
0
1import torch2from typing import Optional3from einops import rearrange4from xfuser.core.distributed import (get_sequence_parallel_rank,5 get_sequence_parallel_world_size,6 get_sp_group)7from xfuser.core.long_ctx_attention import xFuserLongContextAttention8 9def sinusoidal_embedding_1d(dim, position):10 sinusoid = torch.outer(position.type(torch.float64), torch.pow(11 10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))12 x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)13 return x.to(position.dtype)14 15def pad_freqs(original_tensor, target_len):16 seq_len, s1, s2 = original_tensor.shape17 pad_size = target_len - seq_len18 padding_tensor = torch.ones(19 pad_size,20 s1,21 s2,22 dtype=original_tensor.dtype,23 device=original_tensor.device)24 padded_tensor = torch.cat([original_tensor, padding_tensor], dim=0)25 return padded_tensor26 27def rope_apply(x, freqs, num_heads):28 x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)29 s_per_rank = x.shape[1]30 31 x_out = torch.view_as_complex(x.to(torch.float64).reshape(32 x.shape[0], x.shape[1], x.shape[2], -1, 2))33 34 sp_size = get_sequence_parallel_world_size()35 sp_rank = get_sequence_parallel_rank()36 freqs = pad_freqs(freqs, s_per_rank * sp_size)37 freqs_rank = freqs[(sp_rank * s_per_rank):((sp_rank + 1) * s_per_rank), :, :]38 39 x_out = torch.view_as_real(x_out * freqs_rank).flatten(2)40 return x_out.to(x.dtype)41 42def usp_dit_forward(self,43 x: torch.Tensor,44 timestep: torch.Tensor,45 context: torch.Tensor,46 clip_feature: Optional[torch.Tensor] = None,47 y: Optional[torch.Tensor] = None,48 use_gradient_checkpointing: bool = False,49 use_gradient_checkpointing_offload: bool = False,50 **kwargs,51 ):52 t = self.time_embedding(53 sinusoidal_embedding_1d(self.freq_dim, timestep))54 t_mod = self.time_projection(t).unflatten(1, (6, self.dim))55 context = self.text_embedding(context)56 57 if self.has_image_input:58 x = torch.cat([x, y], dim=1) # (b, c_x + c_y, f, h, w)59 clip_embdding = self.img_emb(clip_feature)60 context = torch.cat([clip_embdding, context], dim=1)61 62 x, (f, h, w) = self.patchify(x)63 64 freqs = torch.cat([65 self.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),66 self.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),67 self.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)68 ], dim=-1).reshape(f * h * w, 1, -1).to(x.device)69 70 def create_custom_forward(module):71 def custom_forward(*inputs):72 return module(*inputs)73 return custom_forward74 75 # Context Parallel76 x = torch.chunk(77 x, get_sequence_parallel_world_size(),78 dim=1)[get_sequence_parallel_rank()]79 80 for block in self.blocks:81 if self.training and use_gradient_checkpointing:82 if use_gradient_checkpointing_offload:83 with torch.autograd.graph.save_on_cpu():84 x = torch.utils.checkpoint.checkpoint(85 create_custom_forward(block),86 x, context, t_mod, freqs,87 use_reentrant=False,88 )89 else:90 x = torch.utils.checkpoint.checkpoint(91 create_custom_forward(block),92 x, context, t_mod, freqs,93 use_reentrant=False,94 )95 else:96 x = block(x, context, t_mod, freqs)97 98 x = self.head(x, t)99 100 # Context Parallel101 x = get_sp_group().all_gather(x, dim=1)102 103 # unpatchify104 x = self.unpatchify(x, (f, h, w))105 return x106 107 108def usp_attn_forward(self, x, freqs):109 q = self.norm_q(self.q(x))110 k = self.norm_k(self.k(x))111 v = self.v(x)112 113 q = rope_apply(q, freqs, self.num_heads)114 k = rope_apply(k, freqs, self.num_heads)115 q = rearrange(q, "b s (n d) -> b s n d", n=self.num_heads)116 k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads)117 v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads)118 119 x = xFuserLongContextAttention()(120 None,121 query=q,122 key=k,123 value=v,124 )125 x = x.flatten(2)126 127 del q, k, v128 torch.cuda.empty_cache()129 return self.o(x)