CoolFace
Apppublic

davanstrien/deepseek-ocr

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
clip_sdpa.py505 linesDownload Raw Back to deepencoder
1from contextlib import nullcontext2import math3from typing import Optional, Tuple4# from megatron.model import LayerNorm5from easydict import EasyDict as adict6import torch7from torch.nn import functional as F8from torch import nn9from flash_attn import flash_attn_qkvpacked_func, flash_attn_func10# from optimus import flash_attn_func11# from megatron.core import tensor_parallel12# from megatron.core import parallel_state as mpu13# from megatron.core.utils import make_viewless_tensor, divide14# from megatron.model.fused_rms_norm import RMSNorm15# from megatron.model.transformer import (16#     FlashSelfAttention,17#     NoopTransformerLayer,18#     _cfg_to_kwargs,19# )20# from megatron.model.enums import AttnMaskType, AttnType21# from megatron.model.fused_softmax import FusedScaleMaskSoftmax22# from megatron.model.utils import attention_mask_func23 24# from megatron.model.module import MegatronModule25 26# try:27#     from einops import rearrange28# except ImportError:29#     rearrange = None30 31# from flash_attn import flash_attn_varlen_func as flash_attn_unpadded_func32 33# try:34#     # flash attention 2.x35#     from flash_attn import flash_attn_varlen_func as flash_attn_unpadded_func36# except ImportError:37#     try:38#         # flash attention 1.x39#         from flash_attn.flash_attn_interface import flash_attn_unpadded_func40#     except ImportError:41#         flash_attn_unpadded_func = None42 43# try:44#     from flash_attn.flash_attn_interface import flash_attn_unpadded_relative_attention_bias_func45# except ImportError:46#     flash_attn_unpadded_relative_attention_bias_func = None47 48# try:49#     from flash_attn.flash_attn_interface import mask_flash_attn_unpadded_func50# except ImportError:51#     mask_flash_attn_unpadded_func = None52 53 54class LayerNormfp32(torch.nn.LayerNorm):55    """Subclass torch's LayerNorm to handle fp16."""56 57    def forward(self, x: torch.Tensor):58        orig_type = x.dtype59        ret = super().forward(x.type(torch.float32))60        return ret.type(orig_type)61 62 63def get_abs_pos(abs_pos, tgt_size):64    # abs_pos: L, C65    # tgt_size: M66    # return: M, C67 68    # print(tgt_size)69    # print(abs_pos.shape)70    # exit()71    dim = abs_pos.size(-1)72    # print(dim)73    abs_pos_new = abs_pos.squeeze(0)74    cls_token, old_pos_embed = abs_pos_new[:1], abs_pos_new[1:]75 76 77 78    src_size = int(math.sqrt(abs_pos_new.shape[0] - 1))79    tgt_size = int(math.sqrt(tgt_size))80    dtype = abs_pos.dtype81 82    if src_size != tgt_size:83        old_pos_embed = old_pos_embed.view(1, src_size, src_size, dim).permute(0, 3, 1,84                                                                                    2).contiguous()85        old_pos_embed = old_pos_embed.to(torch.float32)86        new_pos_embed = F.interpolate(87            old_pos_embed,88            size=(tgt_size, tgt_size),89            mode='bicubic',90            antialias=True,91            align_corners=False,92        ).to(dtype)93        new_pos_embed = new_pos_embed.permute(0, 2, 3, 1)94        new_pos_embed = new_pos_embed.view(tgt_size * tgt_size, dim)95        vision_pos_embed = torch.cat([cls_token, new_pos_embed], dim=0)96        vision_pos_embed = vision_pos_embed.view(1, tgt_size * tgt_size + 1, dim)97        return vision_pos_embed98    else:99        return abs_pos100 101@torch.jit.script102def quick_gelu(x):103    return x * torch.sigmoid(1.702 * x)104 105 106 107class CLIPVisionEmbeddings(nn.Module):108    def __init__(self, hidden_size=1024, image_size=224, patch_size=14, num_channels=3):109        super().__init__()110        self.embed_dim = hidden_size111        self.image_size = image_size112        self.patch_size = patch_size113 114        self.class_embedding = torch.nn.Parameter(torch.randn(self.embed_dim))115 116        self.patch_embedding = torch.nn.Conv2d(117            in_channels=num_channels,118            out_channels=self.embed_dim,119            kernel_size=self.patch_size,120            stride=self.patch_size,121            bias=False,122        )123 124        self.num_patches = (self.image_size // self.patch_size) ** 2125        self.num_positions = self.num_patches + 1126        self.position_embedding = torch.nn.Embedding(self.num_positions, self.embed_dim)127        self.register_buffer(128            "position_ids", torch.arange(self.num_positions).expand((1, -1))129        )130 131    def forward(self, pixel_values, patch_embeds):132        batch_size = pixel_values.shape[0]133        # patch_embeds = self.patch_embedding(134        #     pixel_values135        # )  # shape = [*, width, grid, grid]136 137 138        if patch_embeds is not None:139            patch_embeds = patch_embeds140            # print(patch_embeds.shape)141        else:142            patch_embeds = self.patch_embedding(pixel_values)  143            # print(111111)144        # shape = [*, width, grid, grid]145        # patch_embeds = patch_embeds.flatten(2).transpose(1, 2)146 147        patch_embeds = patch_embeds.flatten(2).transpose(1, 2)148 149 150        class_embeds = self.class_embedding.expand(batch_size, 1, -1)151        embeddings = torch.cat([class_embeds, patch_embeds], dim=1)152 153        # x = torch.cat([cls_token, x], dim=1)154        embeddings = embeddings + get_abs_pos(self.position_embedding(self.position_ids), embeddings.size(1))155        # embeddings = embeddings + self.position_embedding(self.position_ids)156        return embeddings157 158 159class NoTPFeedForward(nn.Module):160    def __init__(161            self,162            cfg,163            dim: int,164            hidden_dim: int,165    ):166        super().__init__()167 168        self.fc1 = torch.nn.Linear(dim, hidden_dim, bias=True)169        self.fc2 = torch.nn.Linear(hidden_dim, dim, bias=True)170 171    def forward(self, x):172        output = self.fc2(quick_gelu(self.fc1(x)))173        return output174 175 176# from optimus.flash_attn_interface import flash_attn_qkvpacked_func177 178 179# class NoTPAttention(nn.Module):180#     def __init__(self, cfg):181#         super().__init__()182#         self.num_heads = cfg.num_attention_heads183#         self.n_local_heads = cfg.num_attention_heads184#         self.head_dim = cfg.hidden_size // cfg.num_attention_heads185#         self.max_seq_len = cfg.seq_length186#         self.use_flash_attention = cfg.use_flash_attn187 188#         self.qkv_proj = torch.nn.Linear(cfg.hidden_size, cfg.hidden_size * 3, bias=True)189#         self.out_proj = torch.nn.Linear(cfg.hidden_size, cfg.hidden_size, bias=True)190 191#         # self.core_attention = CoreAttention(cfg, AttnType.self_attn)192 193#         self.attn_drop = cfg.attention_dropout194 195#     def forward(196#             self,197#             x: torch.Tensor,198#     ):199#         bsz, seqlen, _ = x.shape200#         xqkv = self.qkv_proj(x)201#         xqkv = xqkv.view(bsz, seqlen, 3, self.num_heads, self.head_dim)202 203#         if self.use_flash_attention:204#             output = flash_attn_qkvpacked_func(xqkv)205#             output = output.view(bsz, seqlen, -1)206#         else:207#             xq, xk, xv = torch.split(xqkv, 1, dim=2)208#             xq = xq.squeeze(2)209#             xk = xk.squeeze(2)210#             xv = xv.squeeze(2)211#             # xq, xk, xv = xqkv[:, :, 0, ...], xqkv[:, :, 1, ...], xqkv[:, :, 2, ...]212 213#             # (B, num_head, S, head_size)214#             xq = xq.permute(0, 2, 1, 3)215#             xk = xk.permute(0, 2, 1, 3)216#             xv = xv.permute(0, 2, 1, 3)217 218#             output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None)219#             utput = output.permute(0, 2, 1, 3).view(bsz, seqlen, -1)220#         output = self.out_proj(output)221#         return output222 223 224# from optimus.flash_attn_interface import flash_attn_qkvpacked_func225 226 227class NoTPAttention(torch.nn.Module):228    def __init__(self, cfg):229        super().__init__()230        self.num_heads = cfg.num_attention_heads231        self.n_local_heads = cfg.num_attention_heads232        self.head_dim = cfg.hidden_size // cfg.num_attention_heads233        self.max_seq_len = cfg.seq_length234        self.use_flash_attention = cfg.use_flash_attn235 236        self.qkv_proj = torch.nn.Linear(cfg.hidden_size, cfg.hidden_size * 3, bias=True)237        self.out_proj = torch.nn.Linear(cfg.hidden_size, cfg.hidden_size, bias=True)238 239        # self.core_attention = CoreAttention(cfg, AttnType.self_attn)240 241        self.attn_drop = cfg.attention_dropout242 243    def forward(244            self,245            x: torch.Tensor,246    ):247        bsz, seqlen, _ = x.shape248        xqkv = self.qkv_proj(x)249        xqkv = xqkv.view(bsz, seqlen, 3, self.num_heads, self.head_dim)250 251        if self.use_flash_attention:252            output = flash_attn_qkvpacked_func(xqkv)253            output = output.view(bsz, seqlen, -1)254            # xq, xk, xv = torch.split(xqkv, 1, dim=2)255            # xq = xq.squeeze(2)256            # xk = xk.squeeze(2)257            # xv = xv.squeeze(2)258            # # xq, xk, xv = xqkv[:, :, 0, ...], xqkv[:, :, 1, ...], xqkv[:, :, 2, ...]259 260            # # (B, num_head, S, head_size)261            # xq = xq.permute(0, 2, 1, 3)262            # xk = xk.permute(0, 2, 1, 3)263            # xv = xv.permute(0, 2, 1, 3)264            # # with torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=False, enable_mem_efficient=False):265            # output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None)266            # output = output.permute(0, 2, 1, 3).reshape(bsz, seqlen, -1)267                # output = output.permute(0, 2, 1, 3).contiguous().view(bsz, seqlen, -1)268        else:269            # output = flash_attn_qkvpacked_func(xqkv)270            xq, xk, xv = torch.split(xqkv, 1, dim=2)271            xq = xq.squeeze(2)272            xk = xk.squeeze(2)273            xv = xv.squeeze(2)274            # xq, xk, xv = xqkv[:, :, 0, ...], xqkv[:, :, 1, ...], xqkv[:, :, 2, ...]275 276            # (B, num_head, S, head_size)277            xq = xq.permute(0, 2, 1, 3)278            xk = xk.permute(0, 2, 1, 3)279            xv = xv.permute(0, 2, 1, 3)280            # with torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=False, enable_mem_efficient=False):281            output = torch.nn.functional.scaled_dot_product_attention(xq, xk, xv, attn_mask=None)282            output = output.permute(0, 2, 1, 3).reshape(bsz, seqlen, -1)283        output = self.out_proj(output)284        return output285 286class NoTPTransformerBlock(nn.Module):287    def __init__(self, cfg, layer_id: int, multiple_of=256):288        super().__init__()289 290        self.n_heads = cfg.num_attention_heads291        self.dim = cfg.hidden_size292        self.head_dim = cfg.hidden_size // cfg.num_attention_heads293        self.self_attn = NoTPAttention(cfg)294        self.mlp = NoTPFeedForward(295            cfg, dim=cfg.hidden_size, hidden_dim=cfg.ffn_hidden_size296        )297        self.layer_id = layer_id298        self.layer_norm1 = torch.nn.LayerNorm(299            cfg.hidden_size, eps=cfg.layernorm_epsilon300        )301        self.layer_norm2 = torch.nn.LayerNorm(302            cfg.hidden_size, eps=cfg.layernorm_epsilon303        )304 305    def forward(self, x: torch.Tensor):306        residual = self.self_attn.forward(self.layer_norm1(x))307        h = x + residual308        out = h + self.mlp.forward(self.layer_norm2(h))309        return out310 311 312class NoTPTransformer(nn.Module):313    def __init__(self, cfg):314        super().__init__()315 316        self.cfg = cfg317        # self.recompute_list = self.cfg.get("recompute_list", [])318        self.num_layers = cfg.num_layers  # _get_num_layers(cfg)319 320        self.layers = torch.nn.ModuleList()321        for layer_id in range(self.num_layers):322            self.layers.append(323                NoTPTransformerBlock(324                    cfg,325                    layer_id + 1,326                )327            )328 329    def forward(330            self,331            hidden_states,332    ):333 334        for lid, layer in enumerate(self.layers):335            # if lid in self.recompute_list:336            #     def custom(layer_id):337            #         def custom_forward(*args, **kwargs):338            #             x_ = self.layers[layer_id](*args, **kwargs)339            #             return x_340 341            #         return custom_forward342 343            #     assert hidden_states.requires_grad == True, logger.warning(344            #         "When using recalculation, the input must have grad fn"345            #     )346            #     hidden_states = tensor_parallel.checkpoint(347            #         custom(lid),348            #         False,349            #         hidden_states.contiguous()350            #     )351            # else:352            hidden_states = layer(hidden_states)353 354        return hidden_states355 356 357# from megatron.core.tensor_parallel.layers import non_tensor_paralleled, local_dp_reduce, local_dp_scatter358 359class VitModel(nn.Module):360    def __init__(361            self,362            cfg,363            freeze_embed=False,364            freeze_pre_norm=False365    ) -> None:366        super().__init__()367 368        self.embeddings = CLIPVisionEmbeddings(hidden_size=cfg.hidden_size, image_size=cfg.image_size, patch_size=cfg.patch_size)369 370        if freeze_embed:371            for name, param in self.embeddings.named_parameters():372                param.requires_grad = False373 374        self.transformer = NoTPTransformer(cfg=cfg)375 376        if cfg.get("fp32norm", False):377            logger.info("Load fp32 layernorm for ViT.")378            self.pre_layrnorm = LayerNormfp32(379                cfg.hidden_size,380                eps=cfg.get("pre_layernorm_epsilon", 1e-5),381            )382        else:383            self.pre_layrnorm = torch.nn.LayerNorm(384                cfg.hidden_size,385                eps=cfg.get("pre_layernorm_epsilon", 1e-5),386            )387 388        # self.pre_layrnorm = RMSNorm(389        #     cfg.hidden_size,390        #     eps=cfg.get("pre_layernorm_epsilon", 1e-5),391        #     sequence_parallel=False,392        #     use_fp32=True,393        #     use_optimus=True,394        # )395 396        if freeze_pre_norm:397            for name, param in self.pre_layrnorm.named_parameters():398                param.requires_grad = False399 400        for p in self.parameters():401            p.micro_dp = True402 403    def set_input_tensor(self, input_tensor):404        if not isinstance(input_tensor, list):405            input_tensor = [input_tensor]406        self.transformer.set_input_tensor(input_tensor[0])407 408    def __str__(self) -> str:409        return "open_clip"410 411    def forward(412            self,413            x,414            patch_embeds415    ):416        x = self.embeddings(x, patch_embeds)417        hidden_states = self.pre_layrnorm(x)418 419        # hidden_states, dis = local_dp_scatter(hidden_states)420        output = self.transformer(hidden_states)421 422        # output = local_dp_reduce(output, dis)423 424        return output425 426 427vit_model_cfg = adict(428    num_layers=24,429    hidden_size=1024,430    num_heads = 16,431    num_attention_heads=16,432    ffn_hidden_size=4096,433    seq_length=256,434    max_position_embeddings=256,435    use_flash_attn=False,436    understand_projector_stride=2,437    hidden_dropout = 0.0,438    attention_dropout = 0.0,439    no_persist_layer_norm = False,440    layernorm_epsilon = 1e-5,441    pre_layernorm_epsilon = 1e-5,442    image_size = 224,443    patch_size = 14,444    recompute_list = []445)446 447def build_clip_l():448    return VitModel(449        cfg=vit_model_cfg,450        freeze_embed=False,451        freeze_pre_norm=False,452    )453 454 455if __name__ == '__main__':456 457    458    from mmgpt.model.vision_encoder.sam_b import build_sam_vit_b459 460 461 462    vit_model_cfg = adict(463        num_layers=24,464        hidden_size=1024,465        num_attention_heads=16,466        ffn_hidden_size=4096,467        seq_length=256,468        max_position_embeddings=256,469        use_flash_attn=False,470        understand_projector_stride=2,471        hidden_dropout = 0.0,472        attention_dropout = 0.0,473        no_persist_layer_norm = False,474        layernorm_epsilon = 1e-5,475        pre_layernorm_epsilon = 1e-5,476        image_size = 224,477        patch_size = 14,478        recompute_list = []479    )480 481    sam_model = build_sam_vit_b()482 483 484    vision_model = VitModel(485        cfg=vit_model_cfg,486        freeze_embed=False,487        freeze_pre_norm=False,488    )489 490    # model = VitModel(1344)491    # x = torch.zeros(2, 3, 224, 224)492    x = torch.zeros(2, 3, 1024, 1024)493 494    495    with torch.no_grad():496        # y = vision_model(x)497        patch_embed = sam_model(x)498        print(patch_embed.shape)499        y = vision_model(x, patch_embed)500        print(y.shape)501 502        image_feature = torch.add(y[:, 1:], patch_embed.flatten(2).permute(0, 2, 1))503 504        print(image_feature.shape)505