Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
xdit_context_parallel.py129 linesDownload Raw Back to distributed
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)