{"behavior":"degenerate_repetition","label_filter":null,"limit":10,"next_offset":38508,"offset":38498,"previous_offset":null,"rows":[{"assistant_initiates":false,"asst_text":"Hello! I'm just a bunch of algorithms, so I don't have feelings, but I'm here and ready to help you. How can I assist you today?","behavior":"degenerate_repetition","conversation_hash":"d69e9258913f0257f5bc122aed89139b","conversation_redacted":false,"conversation_toxic":false,"country":"United States","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"English","model":"gpt-4o-2024-08-06","n_chars":145,"n_turns_in_conversation":1,"n_words":30,"source_row_offset":38498,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"e2cd81b49bdd1fcec7c55c0a","turn_index":0,"user_text":"hi, how are you?"},{"assistant_initiates":false,"asst_text":"\u5f53\u7136\u53ef\u4ee5\uff01\u8bf7\u95ee\u4f60\u6709\u4ec0\u4e48\u95ee\u9898\uff1f\u6211\u4f1a\u5c3d\u529b\u5e2e\u52a9\u4f60\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":30,"n_turns_in_conversation":22,"n_words":2,"source_row_offset":38499,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"2e09f585be34da6c0795991c","turn_index":0,"user_text":"\u80fd\u56de\u7b54\u95ee\u9898\u5417\uff1f"},{"assistant_initiates":false,"asst_text":"\u5728 `VisionMamba` \u7c7b\u7684 `forward_features` \u65b9\u6cd5\u4e2d\uff0c`Mamba` \u7c7b\u7684\u5b9e\u4f8b\u5728\u89c6\u89c9\u8f6c\u6362\u5668\u67b6\u6784\u4e2d\u626e\u6f14\u4e86\u6df7\u5408\u5c42\u6216\u6a21\u5757\u7684\u89d2\u8272\u3002\u5176\u4e3b\u8981\u4f5c\u7528\u662f\u5bf9\u8f93\u5165\u7279\u5f81\u8fdb\u884c\u53d8\u6362\uff0c\u4fc3\u8fdb\u590d\u6742\u7684\u4e0a\u4e0b\u6587\u76f8\u4e92\u4f5c\u7528\u3002\u5177\u4f53\u800c\u8a00\uff0c\u5728 `Mamba` \u6267\u884c\u8fc7\u7a0b\u4e2d\uff0c\u5b83\u4e3b\u8981\u901a\u8fc7\u4ee5\u4e0b\u51e0\u4e2a\u6b65\u9aa4\u6765\u64cd\u4f5c\u7279\u5f81\uff1a\n\n1. **\u6295\u5f71\u64cd\u4f5c**\uff1a`Mamba` \u4f7f\u7528\u7ebf\u6027\u5c42\u5bf9\u8f93\u5165\u7279\u5f81\u8fdb\u884c\u7ef4\u5ea6\u6295\u5f71\u3002\u8fd9\u5305\u62ec `in_proj`\u3001`x_proj`\u3001\u4ee5\u53ca `dt_proj` \u7b49\u7ebf\u6027\u5c42\uff0c\u7528\u4e8e\u6539\u53d8\u7279\u5f81\u7684\u7ef4\u5ea6\uff0c\u6709\u65f6\u5019\u662f\u6269\u5c55\u7279\u5f81\u7a7a\u95f4\u3002\n\n2. **\u5377\u79ef\u4e0e\u6fc0\u6d3b**\uff1a`Mamba` \u67b6\u6784\u4e2d\u542b\u6709\u4e00\u7ef4\u5377\u79ef\u5c42\uff08`conv1d`\uff09\uff0c\u8be5\u5c42\u53ef\u4ee5\u6355\u6349\u5c40\u90e8\u90bb\u63a5\u4fe1\u606f\uff0c\u5e76\u901a\u8fc7\u6fc0\u6d3b\u51fd\u6570\uff08\u5982 `SiLU`\uff09\u8fdb\u884c\u975e\u7ebf\u6027\u53d8\u6362\u3002\n\n3. **\u72b6\u6001\u7a7a\u95f4\u6a21\u578b\uff08SSM\uff09**\uff1a\u6838\u5fc3\u7684\u5904\u7406\u5728\u4e8e `selective_scan_fn` \u51fd\u6570\uff08\u5047\u5b9a\u5176\u88ab\u6210\u529f\u5bfc\u5165\u53ca\u4f7f\u7528\uff09\uff0c\u5b83\u4f7f\u7528\u72b6\u6001\u7a7a\u95f4\u6a21\u578b\u6280\u672f\u5bf9\u5e8f\u5217\u4fe1\u53f7\u8fdb\u884c\u9ad8\u6548\u5904\u7406\u3002\u8fd9\u90e8\u5206\u64cd\u4f5c\u7279\u522b\u9002\u5408\u957f\u5e8f\u5217\u7684\u6709\u6548\u4fe1\u606f\u6df7\u5408\u4e0e\u901a\u8def\u5b9a\u4e49\u3002\n\n4. **\u8f93\u51fa\u6295\u5f71**\uff1a\u5904\u7406\u5b8c\u6210\u540e\uff0c\u901a\u8fc7 `out_proj` \u5c42\u5c06\u7ed3\u679c\u7279\u5f81\u6620\u5c04\u56de\u5230\u6a21\u578b\u7684\u539f\u59cb\u7ef4\u5ea6\uff0c\u4f7f\u5f97\u540e\u7eed\u5904\u7406\u6216\u8f93\u51fa\u6210\u4e3a\u53ef\u80fd\u3002\n\n\u5728 `VisionMamba` \u7684 `forward_features` \u65b9\u6cd5\u4e2d\uff0c`Mamba` \u7c7b\u7684\u4f5c\u7528\u662f\u63d0\u4f9b\u57fa\u4e8e\u72b6\u6001\u7a7a\u95f4\u7684\u5904\u7406\u65b9\u6cd5\uff0c\u4f7f\u5f97\u6a21\u578b\u53ef\u4ee5\u6709\u6548\u5730\u7ed3\u5408\u957f\u8ddd\u4e0a\u4e0b\u6587\u548c\u5c40\u90e8\u7279\u5f81\uff0c\u4ece\u800c\u5e2e\u52a9\u6a21\u578b\u66f4\u597d\u5730\u8bc6\u522b\u548c\u5206\u6790\u8f93\u5165\u56fe\u50cf\u5e8f\u5217\u4e2d\u7684\u590d\u6742\u6a21\u5f0f\u3002\u901a\u8fc7\u8fd9\u79cd\u590d\u6742\u800c\u7ec6\u81f4\u7684\u7279\u5f81\u6df7\u5408\uff0c`Mamba` \u4fc3\u8fdb\u4e86\u56fe\u50cf\u7279\u5f81\u7684\u6df1\u5c42\u6b21\u63d0\u53d6\u548c\u8bed\u4e49\u7406\u89e3\uff0c\u662f\u589e\u5f3a\u6a21\u578b\u6027\u80fd\u7684\u5173\u952e\u6784\u5efa\u6a21\u5757\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":46273,"n_turns_in_conversation":22,"n_words":3595,"source_row_offset":38500,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"ad573c94d094a0fa7f99a603","turn_index":1,"user_text":"\u7ed3\u5408# Copyright (c) 2015-present, Facebook, Inc.\n# All rights reserved.\nimport torch\nimport torch.nn as nn\nfrom functools import partial\nfrom torch import Tensor\nfrom typing import Optional\n\nfrom timm.models.vision_transformer import VisionTransformer, _cfg\nfrom timm.models.registry import register_model\nfrom timm.models.layers import trunc_normal_, lecun_normal_\n\nfrom timm.models.layers import DropPath, to_2tuple\nfrom timm.models.vision_transformer import _load_weights\n\nimport math\n\nfrom collections import namedtuple\n\nfrom mamba_ssm.modules.mamba_simple import Mamba\nfrom mamba_ssm.utils.generation import GenerationMixin\nfrom mamba_ssm.utils.hf import load_config_hf, load_state_dict_hf\n\nfrom lib.models.mamba_fetrack.rope import *\nimport random\nfrom lib.models.layers.head import build_box_head\nimport importlib\nimport lib.train.admin.settings as ws_settings\nfrom .utils import combine_tokens, recover_tokens\n\ntry:\n    from mamba_ssm.ops.triton.layernorm import RMSNorm, layer_norm_fn, rms_norm_fn\nexcept ImportError:\n    RMSNorm, layer_norm_fn, rms_norm_fn = None, None, None\n    \n\n    \n    \n\n\n__all__ = [\n    'vim_tiny_patch16_224', 'vim_small_patch16_224', 'vim_base_patch16_224',\n    'vim_tiny_patch16_384', 'vim_small_patch16_384', 'vim_base_patch16_384',\n]\n\n\nclass PatchEmbed(nn.Module):\n    \"\"\" 2D Image to Patch Embedding\n    \"\"\"\n    def __init__(self, img_size=224, patch_size=16, stride=16, in_chans=3, embed_dim=768, norm_layer=None, flatten=True):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.grid_size = ((img_size[0] - patch_size[0]) // stride + 1, (img_size[1] - patch_size[1]) // stride + 1)\n        self.num_patches = self.grid_size[0] * self.grid_size[1]\n        self.flatten = flatten\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=stride)\n        self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        # assert H == self.img_size[0] and W == self.img_size[1], \\\n        #     f\"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]}).\"\n        x = self.proj(x)\n        if self.flatten:\n            x = x.flatten(2).transpose(1, 2)  # BCHW -> BNC\n        x = self.norm(x)\n        return x\n    \n\nclass Block(nn.Module):\n    def __init__(\n        self, dim, mixer_cls, norm_cls=nn.LayerNorm, fused_add_norm=False, residual_in_fp32=False,drop_path=0.,\n    ):\n        \"\"\"\n        Simple block wrapping a mixer class with LayerNorm/RMSNorm and residual connection\"\n\n        This Block has a slightly different structure compared to a regular\n        prenorm Transformer block.\n        The standard block is: LN -> MHA/MLP -> Add.\n        [Ref: https://arxiv.org/abs/2002.04745]\n        Here we have: Add -> LN -> Mixer, returning both\n        the hidden_states (output of the mixer) and the residual.\n        This is purely for performance reasons, as we can fuse add and LayerNorm.\n        The residual needs to be provided (except for the very first block).\n        \"\"\"\n        super().__init__()\n        self.residual_in_fp32 = residual_in_fp32\n        self.fused_add_norm = fused_add_norm\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        if self.fused_add_norm:\n            assert RMSNorm is not None, \"RMSNorm import fails\"\n            assert isinstance(\n                self.norm, (nn.LayerNorm, RMSNorm)\n            ), \"Only LayerNorm and RMSNorm are supported for fused_add_norm\"\n\n    def forward(\n        self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None\n    ):\n        r\"\"\"Pass the input through the encoder layer.\n\n        Args:\n            hidden_states: the sequence to the encoder layer (required).\n            residual: hidden_states = Mixer(LN(residual))\n        \"\"\"\n        if not self.fused_add_norm:\n            if residual is None:\n                residual = hidden_states\n            else:\n                residual = residual + self.drop_path(hidden_states)\n            \n            hidden_states = self.norm(residual.to(dtype=self.norm.weight.dtype))\n            if self.residual_in_fp32:\n                residual = residual.to(torch.float32)\n        else:\n            fused_add_norm_fn = rms_norm_fn if isinstance(self.norm, RMSNorm) else layer_norm_fn\n            if residual is None:\n                hidden_states, residual = fused_add_norm_fn(\n                    hidden_states,\n                    self.norm.weight,\n                    self.norm.bias,\n                    residual=residual,\n                    prenorm=True,\n                    residual_in_fp32=self.residual_in_fp32,\n                    eps=self.norm.eps,\n                )\n            else:\n                hidden_states, residual = fused_add_norm_fn(\n                    self.drop_path(hidden_states),\n                    self.norm.weight,\n                    self.norm.bias,\n                    residual=residual,\n                    prenorm=True,\n                    residual_in_fp32=self.residual_in_fp32,\n                    eps=self.norm.eps,\n                )    \n        hidden_states = self.mixer(hidden_states, inference_params=inference_params)\n        return hidden_states, residual\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        return self.mixer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype, **kwargs)\n\n\ndef create_block(\n    d_model,\n    ssm_cfg=None,\n    norm_epsilon=1e-5,\n    drop_path=0.,\n    rms_norm=False,\n    residual_in_fp32=False,\n    fused_add_norm=False,\n    layer_idx=None,\n    device=None,\n    dtype=None,\n    if_bimamba=False,\n    bimamba_type=\"none\",\n    if_devide_out=False,\n    init_layer_scale=None,\n):\n    if if_bimamba:\n        bimamba_type = \"v1\"\n    if ssm_cfg is None:\n        ssm_cfg = {}\n    factory_kwargs = {\"device\": device, \"dtype\": dtype}\n    mixer_cls = partial(Mamba, layer_idx=layer_idx, bimamba_type=bimamba_type, if_devide_out=if_devide_out, init_layer_scale=init_layer_scale, **ssm_cfg, **factory_kwargs)\n    norm_cls = partial(\n        nn.LayerNorm if not rms_norm else RMSNorm, eps=norm_epsilon, **factory_kwargs\n    )\n    block = Block(\n        d_model,\n        mixer_cls,\n        norm_cls=norm_cls,\n        drop_path=drop_path,\n        fused_add_norm=fused_add_norm,\n        residual_in_fp32=residual_in_fp32,\n    )\n    block.layer_idx = layer_idx\n    return block\n\n\n# https://github.com/huggingface/transformers/blob/c28d04e9e252a1a099944e325685f14d242ecdcd/src/transformers/models/gpt2/modeling_gpt2.py#L454\ndef _init_weights(\n    module,\n    n_layer,\n    initializer_range=0.02,  # Now only used for embedding layer.\n    rescale_prenorm_residual=True,\n    n_residuals_per_layer=1,  # Change to 2 if we have MLP\n):\n    if isinstance(module, nn.Linear):\n        if module.bias is not None:\n            if not getattr(module.bias, \"_no_reinit\", False):\n                nn.init.zeros_(module.bias)\n    elif isinstance(module, nn.Embedding):\n        nn.init.normal_(module.weight, std=initializer_range)\n\n    if rescale_prenorm_residual:\n        # Reinitialize selected weights subject to the OpenAI GPT-2 Paper Scheme:\n        #   > A modified initialization which accounts for the accumulation on the residual path with model depth. Scale\n        #   > the weights of residual layers at initialization by a factor of 1/\u221aN where N is the # of residual layers.\n        #   >   -- GPT-2 :: https://openai.com/blog/better-language-models/\n        #\n        # Reference (Megatron-LM): https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/model/gpt_model.py\n        for name, p in module.named_parameters():\n            if name in [\"out_proj.weight\", \"fc2.weight\"]:\n                # Special Scaled Initialization --> There are 2 Layer Norms per Transformer Block\n                # Following Pytorch init, except scale by 1/sqrt(2 * n_layer)\n                # We need to reinit p since this code could be called multiple times\n                # Having just p *= scale would repeatedly scale it down\n                nn.init.kaiming_uniform_(p, a=math.sqrt(5))\n                with torch.no_grad():\n                    p /= math.sqrt(n_residuals_per_layer * n_layer)\n\n\ndef segm_init_weights(m):\n    if isinstance(m, nn.Linear):\n        trunc_normal_(m.weight, std=0.02)\n        if isinstance(m, nn.Linear) and m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n    elif isinstance(m, nn.Conv2d):\n        # NOTE conv was left to pytorch default in my original init\n        lecun_normal_(m.weight)\n        if m.bias is not None:\n            nn.init.zeros_(m.bias)\n    elif isinstance(m, (nn.LayerNorm, nn.GroupNorm, nn.BatchNorm2d)):\n        nn.init.zeros_(m.bias)\n        nn.init.ones_(m.weight)\n\n\nclass VisionMamba(nn.Module):\n    def __init__(self, \n                 img_size=224, \n                 patch_size=16, \n                 stride=16,\n                 depth=24, \n                 embed_dim=192, \n                 channels=3, \n                 num_classes=1000,\n                 ssm_cfg=None, \n                 drop_rate=0.,\n                 drop_path_rate=0.1,\n                 norm_epsilon: float = 1e-5, \n                 rms_norm: bool = False, \n                 initializer_cfg=None,\n                 fused_add_norm=False,\n                 residual_in_fp32=False,\n                 device=None,\n                 dtype=None,\n                 ft_seq_len=None,\n                 pt_hw_seq_len=14,\n                 if_bidirectional=False,\n                 final_pool_type='none',\n                 if_abs_pos_embed=False,\n                 if_rope=False,\n                 if_rope_residual=False,\n                 flip_img_sequences_ratio=-1.,\n                 if_bimamba=False,\n                 bimamba_type=\"none\",\n                 if_cls_token=False,\n                 if_devide_out=False,\n                 init_layer_scale=None,\n                 use_double_cls_token=False,\n                 use_middle_cls_token=False,\n                 **kwargs):\n        factory_kwargs = {\"device\": device, \"dtype\": dtype}\n        # add factory_kwargs into kwargs\n        kwargs.update(factory_kwargs) \n        super().__init__()\n        self.residual_in_fp32 = residual_in_fp32\n        self.fused_add_norm = fused_add_norm\n        self.if_bidirectional = if_bidirectional\n        self.final_pool_type = final_pool_type\n        self.if_abs_pos_embed = if_abs_pos_embed\n        self.if_rope = if_rope\n        self.if_rope_residual = if_rope_residual\n        self.flip_img_sequences_ratio = flip_img_sequences_ratio\n        # self.if_cls_token = if_cls_token\n        self.if_cls_token = False\n        self.use_double_cls_token = use_double_cls_token\n        self.use_middle_cls_token = use_middle_cls_token\n        self.num_tokens = 1 if if_cls_token else 0\n\n        # pretrain parameters\n        self.num_classes = num_classes\n        self.d_model = self.num_features = self.embed_dim = embed_dim  # num_features for consistency with other models\n\n        self.patch_embed = PatchEmbed(\n            img_size=img_size, patch_size=patch_size, stride=stride, in_chans=channels, embed_dim=embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        if if_cls_token:\n            if use_double_cls_token:\n                self.cls_token_head = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                self.cls_token_tail = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                self.num_tokens = 2\n            else:\n                self.cls_token = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                # self.num_tokens = 1\n            \n        if if_abs_pos_embed:\n            # self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + self.num_tokens, self.embed_dim))\n            self.pos_embed_x = nn.Parameter(torch.zeros(1, 256, self.embed_dim))\n            self.pos_embed_z = nn.Parameter(torch.zeros(1, 64, self.embed_dim))\n            self.pos_drop = nn.Dropout(p=drop_rate)\n\n        if if_rope:\n            half_head_dim = embed_dim // 2\n            hw_seq_len = img_size // patch_size\n            self.rope = VisionRotaryEmbeddingFast(\n                dim=half_head_dim,\n                pt_seq_len=pt_hw_seq_len,\n                ft_seq_len=hw_seq_len\n            )\n        # self.head = nn.Linear(self.num_features, num_classes) if num_classes > 0 else nn.Identity()\n        # settings = ws_settings.Settings()\n        # config_module = importlib.import_module(\"lib.config.%s.config\" % settings.script_name)\n        # cfg = config_module.cfg\n        # self.head = build_box_head(cfg, embed_dim)\n\n\n        # TODO: release this comment\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]  # stochastic depth decay rule\n        # import ipdb;ipdb.set_trace()\n        inter_dpr = [0.0] + dpr\n        self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()\n                # transformer blocks\n        self.layers = nn.ModuleList(\n            [\n                create_block(\n                    embed_dim,\n                    ssm_cfg=ssm_cfg,\n                    norm_epsilon=norm_epsilon,\n                    rms_norm=rms_norm,\n                    residual_in_fp32=residual_in_fp32,\n                    fused_add_norm=fused_add_norm,\n                    layer_idx=i,\n                    if_bimamba=if_bimamba,\n                    bimamba_type=bimamba_type,\n                    drop_path=inter_dpr[i],\n                    if_devide_out=if_devide_out,\n                    init_layer_scale=init_layer_scale,\n                    **factory_kwargs,\n                )\n                for i in range(depth)\n            ]\n        )\n        \n        # output head\n        self.norm_f = (nn.LayerNorm if not rms_norm else RMSNorm)(\n            embed_dim, eps=norm_epsilon, **factory_kwargs\n        )\n\n        # self.pre_logits = nn.Identity()\n\n        # original init\n        self.patch_embed.apply(segm_init_weights)\n        # self.head.apply(segm_init_weights)\n        if if_abs_pos_embed:\n            trunc_normal_(self.pos_embed_x, std=.02)\n            trunc_normal_(self.pos_embed_z, std=.02)\n            \n        if if_cls_token:\n            if use_double_cls_token:\n                trunc_normal_(self.cls_token_head, std=.02)\n                trunc_normal_(self.cls_token_tail, std=.02)\n            else:\n                trunc_normal_(self.cls_token, std=.02)\n\n        # mamba init\n        self.apply(\n            partial(\n                _init_weights,\n                n_layer=depth,\n                **(initializer_cfg if initializer_cfg is not None else {}),\n            )\n        )\n\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        return {\n            i: layer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype, **kwargs)\n            for i, layer in enumerate(self.layers)\n        }\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {\"pos_embed\", \"cls_token\", \"dist_token\", \"cls_token_head\", \"cls_token_tail\"}\n\n    @torch.jit.ignore()\n    def load_pretrained(self, checkpoint_path, prefix=\"\"):\n        _load_weights(self, checkpoint_path, prefix)\n\n    def forward_features(self, z, x, inference_params=None, if_random_cls_token_position=False, if_random_token_rank=False):\n        # taken from https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/vision_transformer.py\n        # with slight modifications to add the dist_token\n        x = self.patch_embed(x)                  #x.shape = torch.Size([B, 3, 256, 256])  -> torch.Size([2, 256, 384])\n        z = self.patch_embed(z)                  #z.shape = torch.Size([B, 3, 128, 128])  -> torch.Size([2, 64, 384])\n        B, M, _ = x.shape\n       \n        if self.if_cls_token:                 # False\n            if self.use_double_cls_token:\n                cls_token_head = self.cls_token_head.expand(B, -1, -1)\n                cls_token_tail = self.cls_token_tail.expand(B, -1, -1)\n                token_position = [0, M + 1]\n                x = torch.cat((cls_token_head, x, cls_token_tail), dim=1)\n                M = x.shape[1]\n            else:\n                if self.use_middle_cls_token:\n                    cls_token = self.cls_token.expand(B, -1, -1)\n                    token_position = M // 2\n                    # add cls token in the middle\n                    x = torch.cat((x[:, :token_position, :], cls_token, x[:, token_position:, :]), dim=1)       \n                elif if_random_cls_token_position:\n                    cls_token = self.cls_token.expand(B, -1, -1)\n                    token_position = random.randint(0, M)\n                    x = torch.cat((x[:, :token_position, :], cls_token, x[:, token_position:, :]), dim=1)\n                    print(\"token_position: \", token_position)\n                else:\n                    cls_token = self.cls_token.expand(B, -1, -1)  # stole cls_tokens impl from Phil Wang, thanks\n                    token_position = 0\n                    x = torch.cat((cls_token, x), dim=1)\n                M = x.shape[1]                 \n       \n        if self.if_abs_pos_embed:                  # True \n            x = x + self.pos_embed_x               # x = x + positon_embemding =torch.Size([B, 256, 384]) + torch.Size([1, 256, 384]) = torch.Size([B, 256, 384])\n            z = z + self.pos_embed_z               # z = z + positon_embemding =torch.Size([B, 64, 384]) + torch.Size([1, 64, 384]) = torch.Size([B, 64, 384])\n            x = torch.cat((z, x), dim=1)           # torch.Size([B, 320, 384])\n            x = self.pos_drop(x)                   # x.shape = torch.Size([B, 320, 384])\n            \n        if if_random_token_rank:                   #False\n            # \u751f\u6210\u968f\u673a shuffle \u7d22\u5f15\n            shuffle_indices = torch.randperm(M)\n\n            if isinstance(token_position, list):\n                print(\"original value: \", x[0, token_position[0], 0], x[0, token_position[1], 0])\n            else:\n                print(\"original value: \", x[0, token_position, 0])\n            print(\"original token_position: \", token_position)\n\n            # \u6267\u884c shuffle\n            x = x[:, shuffle_indices, :]\n\n            if isinstance(token_position, list):\n                # \u627e\u5230 cls token \u5728 shuffle \u4e4b\u540e\u7684\u65b0\u4f4d\u7f6e\n                new_token_position = [torch.where(shuffle_indices == token_position[i])[0].item() for i in range(len(token_position))]\n                token_position = new_token_position\n            else:\n                # \u627e\u5230 cls token \u5728 shuffle \u4e4b\u540e\u7684\u65b0\u4f4d\u7f6e\n                token_position = torch.where(shuffle_indices == token_position)[0].item()\n\n            if isinstance(token_position, list):\n                print(\"new value: \", x[0, token_position[0], 0], x[0, token_position[1], 0])\n            else:\n                print(\"new value: \", x[0, token_position, 0])\n            print(\"new token_position: \", token_position)\n\n\n        if_flip_img_sequences = False\n        if self.flip_img_sequences_ratio > 0 and (self.flip_img_sequences_ratio - random.random()) > 1e-5:        # False\n            x = x.flip([1])\n            if_flip_img_sequences = True\n\n        # mamba impl\n        residual = None\n        hidden_states = x\n        if not self.if_bidirectional:                                 # True\n            for layer in self.layers:\n                if if_flip_img_sequences and self.if_rope:            # False\n                    hidden_states = hidden_states.flip([1])\n                    if residual is not None:\n                        residual = residual.flip([1])\n\n                # rope about\n                if self.if_rope:                                       # False\n                    hidden_states = self.rope(hidden_states)\n                    if residual is not None and self.if_rope_residual:\n                        residual = self.rope(residual)\n\n                if if_flip_img_sequences and self.if_rope:             # False\n                    hidden_states = hidden_states.flip([1])\n                    if residual is not None:\n                        residual = residual.flip([1])\n\n                hidden_states, residual = layer(\n                    hidden_states, residual, inference_params=inference_params\n                )\n        \n        else:             # False\n            # get two layers in a single for-loop\n            for i in range(len(self.layers) // 2):\n                if self.if_rope:\n                    hidden_states = self.rope(hidden_states)\n                    if residual is not None and self.if_rope_residual:\n                        residual = self.rope(residual)\n\n                hidden_states_f, residual_f = self.layers[i * 2](\n                    hidden_states, residual, inference_params=inference_params\n                )\n                hidden_states_b, residual_b = self.layers[i * 2 + 1](\n                    hidden_states.flip([1]), None if residual == None else residual.flip([1]), inference_params=inference_params\n                )\n                hidden_states = hidden_states_f + hidden_states_b.flip([1])\n                residual = residual_f + residual_b.flip([1])\n      \n        if not self.fused_add_norm:         #False\n            if residual is None:\n                residual = hidden_states\n            else:\n                residual = residual + self.drop_path(hidden_states)\n            hidden_states = self.norm_f(residual.to(dtype=self.norm_f.weight.dtype))\n        else:       #True\n            # Set prenorm=False here since we don't need the residual\n            fused_add_norm_fn = rms_norm_fn if isinstance(self.norm_f, RMSNorm) else layer_norm_fn\n            hidden_states = fused_add_norm_fn(                                         # hidden_states.shape = torch.Size([B, 320, 384])\n                self.drop_path(hidden_states),\n                self.norm_f.weight,\n                self.norm_f.bias,\n                eps=self.norm_f.eps,\n                residual=residual,\n                prenorm=False,\n                residual_in_fp32=self.residual_in_fp32,\n            )\n\n        # return only cls token if it exists\n        if self.if_cls_token:          #False\n            if self.use_double_cls_token:\n                return (hidden_states[:, token_position[0], :] + hidden_states[:, token_position[1], :]) / 2\n            else:\n                if self.use_middle_cls_token:\n                    return hidden_states[:, token_position, :]\n                elif if_random_cls_token_position:\n                    return hidden_states[:, token_position, :]\n                else:\n                    return hidden_states[:, token_position, :]\n\n        if self.final_pool_type == 'none':\n            return hidden_states[:, -1, :]\n        elif self.final_pool_type == 'mean':         #True\n            return hidden_states.mean(dim=1)         #hidden_states.shape = torch.Size([2, 384])\n        elif self.final_pool_type == 'max':\n            return hidden_states\n        elif self.final_pool_type == 'all':\n            return hidden_states\n        else:\n            raise NotImplementedError\n\n    def forward(self, x, return_features=False, inference_params=None, if_random_cls_token_position=False, if_random_token_rank=False):\n        x = self.forward_features(x, inference_params, if_random_cls_token_position=if_random_cls_token_position, if_random_token_rank=if_random_token_rank)\n        if return_features:\n            return x\n        x = self.head(x)\n        if self.final_pool_type == 'max':\n            x = x.max(dim=1)[0]\n        return x\n\n\n@register_model\ndef vim_tiny_patch16_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, embed_dim=192, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n@register_model\ndef vim_tiny_patch16_stride8_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, stride=8, embed_dim=192, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n@register_model\ndef vim_small_patch16_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, embed_dim=384, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='all', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    \n    if pretrained:\n        # checkpoint = torch.hub.load_state_dict_from_url(\n        #     url=\"to.do\",\n        #     map_location=\"cpu\", check_hash=True\n        # )\n        # model.load_state_dict(checkpoint[\"model\"])\n        checkpoint = torch.load(pretrained, map_location=\"cpu\")\n        missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"model\"], strict=False)\n        print('Load pretrained model from: ' + pretrained)\n    return model\n\n@register_model\ndef vim_small_patch16_stride8_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, stride=8, embed_dim=384, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n\u548c# Copyright (c) 2023, Tri Dao, Albert Gu.\n\nimport math\nfrom typing import Optional\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import Tensor\n\nfrom einops import rearrange, repeat\n\ntry:\n    from causal_conv1d import causal_conv1d_fn, causal_conv1d_update\nexcept ImportError:\n    causal_conv1d_fn, causal_conv1d_update = None\n\ntry:\n    from mamba_ssm.ops.selective_scan_interface import selective_scan_fn, mamba_inner_fn, bimamba_inner_fn, mamba_inner_fn_no_out_proj\nexcept ImportError:\n    selective_scan_fn, mamba_inner_fn, bimamba_inner_fn, mamba_inner_fn_no_out_proj = None, None, None, None, None\n\ntry:\n    from mamba_ssm.ops.triton.selective_state_update import selective_state_update\nexcept ImportError:\n    selective_state_update = None\n\ntry:\n    from mamba_ssm.ops.triton.layernorm import RMSNorm, layer_norm_fn, rms_norm_fn\nexcept ImportError:\n    RMSNorm, layer_norm_fn, rms_norm_fn = None, None, None\n\n# d_state\u53c8\u79f0N\uff0c\u662f\u9690\u72b6\u6001\u7684\u7ef4\u5ea6\u3002d_model\u4ee3\u8868\u6570\u636e\u5728\u7ecf\u8fc7SSM\u5757\u4e4b\u524d\u548c\u4e4b\u540e\u7684\u6570\u636e\u7ef4\u5ea6\uff0c\u662f\u8f93\u5165x\u7684\u7ef4\u5ea6\nclass Mamba(nn.Module):\n    def __init__(\n        self,\n        d_model,\n        d_state=16,\n        d_conv=4,\n        expand=2,\n        dt_rank=\"auto\",\n        dt_min=0.001,\n        dt_max=0.1,\n        dt_init=\"random\",\n        dt_scale=1.0,\n        dt_init_floor=1e-4,\n        conv_bias=True,\n        bias=False,\n        use_fast_path=True,  # Fused kernel options\n        layer_idx=None,\n        device=None,\n        dtype=None,\n        bimamba_type=\"none\",\n        if_devide_out=False,\n        init_layer_scale=None,\n    ):\n        factory_kwargs = {\"device\": device, \"dtype\": dtype}\n        super().__init__()\n        self.d_model = d_model\n        self.d_state = d_state\n        self.d_conv = d_conv\n        self.expand = expand\n        self.d_inner = int(self.expand * self.d_model)\n        self.dt_rank = math.ceil(self.d_model / 16) if dt_rank == \"auto\" else dt_rank\n        self.use_fast_path = use_fast_path\n        self.layer_idx = layer_idx\n        self.bimamba_type = bimamba_type\n        self.if_devide_out = if_devide_out\n\n        self.init_layer_scale = init_layer_scale\n        if init_layer_scale is not None:\n            self.gamma = nn.Parameter(init_layer_scale * torch.ones((d_model)), requires_grad=True)\n\n        self.in_proj = nn.Linear(self.d_model, self.d_inner * 2, bias=bias, **factory_kwargs)\n\n        self.conv1d = nn.Conv1d(\n            in_channels=self.d_inner,\n            out_channels=self.d_inner,\n            bias=conv_bias,\n            kernel_size=d_conv,\n            groups=self.d_inner,\n            padding=d_conv - 1,\n            **factory_kwargs,\n        )\n\n        self.activation = \"silu\"\n        self.act = nn.SiLU()\n\n        self.x_proj = nn.Linear(\n            self.d_inner, self.dt_rank + self.d_state * 2, bias=False, **factory_kwargs\n        )\n        self.dt_proj = nn.Linear(self.dt_rank, self.d_inner, bias=True, **factory_kwargs)\n\n        # Initialize special dt projection to preserve variance at initialization\n        dt_init_std = self.dt_rank**-0.5 * dt_scale\n        if dt_init == \"constant\":\n            nn.init.constant_(self.dt_proj.weight, dt_init_std)\n        elif dt_init == \"random\":\n            nn.init.uniform_(self.dt_proj.weight, -dt_init_std, dt_init_std)\n        else:\n            raise NotImplementedError\n\n        # Initialize dt bias so that F.softplus(dt_bias) is between dt_min and dt_max\n        dt = torch.exp(\n            torch.rand(self.d_inner, **factory_kwargs) * (math.log(dt_max) - math.log(dt_min))\n            + math.log(dt_min)\n        ).clamp(min=dt_init_floor)\n        # Inverse of softplus: https://github.com/pytorch/pytorch/issues/72759\n        inv_dt = dt + torch.log(-torch.expm1(-dt))\n        with torch.no_grad():\n            self.dt_proj.bias.copy_(inv_dt)\n        # Our initialization would set all Linear.bias to zero, need to mark this one as _no_reinit\n        self.dt_proj.bias._no_reinit = True\n\n        # S4D real initialization\n        A = repeat(\n            torch.arange(1, self.d_state + 1, dtype=torch.float32, device=device),\n            \"n -> d n\",\n            d=self.d_inner,\n        ).contiguous()\n        A_log = torch.log(A)  # Keep A_log in fp32\n        self.A_log = nn.Parameter(A_log)\n        self.A_log._no_weight_decay = True\n\n        # D \"skip\" parameter\n        self.D = nn.Parameter(torch.ones(self.d_inner, device=device))  # Keep in fp32\n        self.D._no_weight_decay = True\n\n        # bidirectional\n        if bimamba_type == \"v1\":\n            A_b = repeat(\n                torch.arange(1, self.d_state + 1, dtype=torch.float32, device=device),\n                \"n -> d n\",\n                d=self.d_inner,\n            ).contiguous()\n            A_b_log = torch.log(A_b)  # Keep A_b_log in fp32\n            self.A_b_log = nn.Parameter(A_b_log)\n            self.A_b_log._no_weight_decay = True\n        elif bimamba_type == \"v2\":\n            A_b = repeat(\n                torch.arange(1, self.d_state + 1, dtype=torch.float32, device=device),\n                \"n -> d n\",\n                d=self.d_inner,\n            ).contiguous()\n            A_b_log = torch.log(A_b)  # Keep A_b_log in fp32\n            self.A_b_log = nn.Parameter(A_b_log)\n            self.A_b_log._no_weight_decay = True \n\n            self.conv1d_b = nn.Conv1d(\n                in_channels=self.d_inner,\n                out_channels=self.d_inner,\n                bias=conv_bias,\n                kernel_size=d_conv,\n                groups=self.d_inner,\n                padding=d_conv - 1,\n                **factory_kwargs,\n            )\n\n            self.x_proj_b = nn.Linear(\n                self.d_inner, self.dt_rank + self.d_state * 2, bias=False, **factory_kwargs\n            )\n            self.dt_proj_b = nn.Linear(self.dt_rank, self.d_inner, bias=True, **factory_kwargs)\n\n            self.D_b = nn.Parameter(torch.ones(self.d_inner, device=device))  # Keep in fp32\n            self.D_b._no_weight_decay = True\n\n        self.out_proj = nn.Linear(self.d_inner, self.d_model, bias=bias, **factory_kwargs)\n\n    def forward(self, hidden_states, inference_params=None):\n        \"\"\"\n        hidden_states: (B, L, D)\n        Returns: same shape as hidden_states\n        \"\"\"\n        batch, seqlen, dim = hidden_states.shape\n\n        conv_state, ssm_state = None, None\n        if inference_params is not None:\n            conv_state, ssm_state = self._get_states_from_cache(inference_params, batch)\n            if inference_params.seqlen_offset > 0:\n                # The states are updated inplace\n                out, _, _ = self.step(hidden_states, conv_state, ssm_state)\n                return out\n\n        # We do matmul and transpose BLH -> HBL at the same time\n        xz = rearrange(\n            self.in_proj.weight @ rearrange(hidden_states, \"b l d -> d (b l)\"),\n            \"d (b l) -> b d l\",\n            l=seqlen,\n        )\n        if self.in_proj.bias is not None:\n            xz = xz + rearrange(self.in_proj.bias.to(dtype=xz.dtype), \"d -> d 1\")\n\n        A = -torch.exp(self.A_log.float())  # (d_inner, d_state)\n        # In the backward pass we write dx and dz next to each other to avoid torch.cat\n        if self.use_fast_path and inference_params is None:  # Doesn't support outputting the states\n            if self.bimamba_type == \"v1\":\n                A_b = -torch.exp(self.A_b_log.float())\n                out = bimamba_inner_fn(\n                    xz,\n                    self.conv1d.weight,\n                    self.conv1d.bias,\n                    self.x_proj.weight,\n                    self.dt_proj.weight,\n                    self.out_proj.weight,\n                    self.out_proj.bias,\n                    A,\n                    A_b,\n                    None,  # input-dependent B\n                    None,  # input-dependent C\n                    self.D.float(),\n                    delta_bias=self.dt_proj.bias.float(),\n                    delta_softplus=True,\n                )    \n            elif self.bimamba_type == \"v2\":\n                A_b = -torch.exp(self.A_b_log.float())\n                out = mamba_inner_fn_no_out_proj(\n                    xz,\n                    self.conv1d.weight,\n                    self.conv1d.bias,\n                    self.x_proj.weight,\n                    self.dt_proj.weight,\n                    A,\n                    None,  # input-dependent B\n                    None,  # input-dependent C\n                    self.D.float(),\n                    delta_bias=self.dt_proj.bias.float(),\n                    delta_softplus=True,\n                )\n                out_b = mamba_inner_fn_no_out_proj(\n                    xz.flip([-1]),\n                    self.conv1d_b.weight,\n                    self.conv1d_b.bias,\n                    self.x_proj_b.weight,\n                    self.dt_proj_b.weight,\n                    A_b,\n                    None,\n                    None,\n                    self.D_b.float(),\n                    delta_bias=self.dt_proj_b.bias.float(),\n                    delta_softplus=True,\n                )\n                # F.linear(rearrange(out_z, \"b d l -> b l d\"), out_proj_weight, out_proj_bias)\n                if not self.if_devide_out:\n                    out = F.linear(rearrange(out + out_b.flip([-1]), \"b d l -> b l d\"), self.out_proj.weight, self.out_proj.bias)\n                else:\n                    out = F.linear(rearrange(out + out_b.flip([-1]), \"b d l -> b l d\") / 2, self.out_proj.weight, self.out_proj.bias)\n\n            else:\n                out = mamba_inner_fn(\n                    xz,\n                    self.conv1d.weight,\n                    self.conv1d.bias,\n                    self.x_proj.weight,\n                    self.dt_proj.weight,\n                    self.out_proj.weight,\n                    self.out_proj.bias,\n                    A,\n                    None,  # input-dependent B\n                    None,  # input-dependent C\n                    self.D.float(),\n                    delta_bias=self.dt_proj.bias.float(),\n                    delta_softplus=True,\n                )\n        else:\n            x, z = xz.chunk(2, dim=1)\n            # Compute short convolution\n            if conv_state is not None:\n                # If we just take x[:, :, -self.d_conv :], it will error if seqlen < self.d_conv\n                # Instead F.pad will pad with zeros if seqlen < self.d_conv, and truncate otherwise.\n                conv_state.copy_(F.pad(x, (self.d_conv - x.shape[-1], 0)))  # Update state (B D W)\n            if causal_conv1d_fn is None:\n                x = self.act(self.conv1d(x)[..., :seqlen])\n            else:\n                assert self.activation in [\"silu\", \"swish\"]\n                x = causal_conv1d_fn(\n                    x=x,\n                    weight=rearrange(self.conv1d.weight, \"d 1 w -> d w\"),\n                    bias=self.conv1d.bias,\n                    activation=self.activation,\n                )\n\n            # We're careful here about the layout, to avoid extra transposes.\n            # We want dt to have d as the slowest moving dimension\n            # and L as the fastest moving dimension, since those are what the ssm_scan kernel expects.\n            x_dbl = self.x_proj(rearrange(x, \"b d l -> (b l) d\"))  # (bl d)\n            dt, B, C = torch.split(x_dbl, [self.dt_rank, self.d_state, self.d_state], dim=-1)\n            dt = self.dt_proj.weight @ dt.t()\n            dt = rearrange(dt, \"d (b l) -> b d l\", l=seqlen)\n            B = rearrange(B, \"(b l) dstate -> b dstate l\", l=seqlen).contiguous()\n            C = rearrange(C, \"(b l) dstate -> b dstate l\", l=seqlen).contiguous()\n            assert self.activation in [\"silu\", \"swish\"]\n            y = selective_scan_fn(\n                x,\n                dt,\n                A,\n                B,\n                C,\n                self.D.float(),\n                z=z,\n                delta_bias=self.dt_proj.bias.float(),\n                delta_softplus=True,\n                return_last_state=ssm_state is not None,\n            )\n            if ssm_state is not None:\n                y, last_state = y\n                ssm_state.copy_(last_state)\n            y = rearrange(y, \"b d l -> b l d\")\n            out = self.out_proj(y)\n        if self.init_layer_scale is not None:\n                out = out * self.gamma    \n        return out\n\n    def step(self, hidden_states, conv_state, ssm_state):\n        dtype = hidden_states.dtype\n        assert hidden_states.shape[1] == 1, \"Only support decoding with 1 token at a time for now\"\n        xz = self.in_proj(hidden_states.squeeze(1))  # (B 2D)\n        x, z = xz.chunk(2, dim=-1)  # (B D)\n\n        # Conv step\n        if causal_conv1d_update is None:\n            conv_state.copy_(torch.roll(conv_state, shifts=-1, dims=-1))  # Update state (B D W)\n            conv_state[:, :, -1] = x\n            x = torch.sum(conv_state * rearrange(self.conv1d.weight, \"d 1 w -> d w\"), dim=-1)  # (B D)\n            if self.conv1d.bias is not None:\n                x = x + self.conv1d.bias\n            x = self.act(x).to(dtype=dtype)\n        else:\n            x = causal_conv1d_update(\n                x,\n                conv_state,\n                rearrange(self.conv1d.weight, \"d 1 w -> d w\"),\n                self.conv1d.bias,\n                self.activation,\n            )\n\n        x_db = self.x_proj(x)  # (B dt_rank+2*d_state)\n        dt, B, C = torch.split(x_db, [self.dt_rank, self.d_state, self.d_state], dim=-1)\n        # Don't add dt_bias here\n        dt = F.linear(dt, self.dt_proj.weight)  # (B d_inner)\n        A = -torch.exp(self.A_log.float())  # (d_inner, d_state)\n\n        # SSM step\n        if selective_state_update is None:\n            # Discretize A and B\n            dt = F.softplus(dt + self.dt_proj.bias.to(dtype=dt.dtype))\n            dA = torch.exp(torch.einsum(\"bd,dn->bdn\", dt, A))\n            dB = torch.einsum(\"bd,bn->bdn\", dt, B)\n            ssm_state.copy_(ssm_state * dA + rearrange(x, \"b d -> b d 1\") * dB)\n            y = torch.einsum(\"bdn,bn->bd\", ssm_state.to(dtype), C)\n            y = y + self.D.to(dtype) * x\n            y = y * self.act(z)  # (B D)\n        else:\n            y = selective_state_update(\n                ssm_state, x, dt, A, B, C, self.D, z=z, dt_bias=self.dt_proj.bias, dt_softplus=True\n            )\n\n        out = self.out_proj(y)\n        return out.unsqueeze(1), conv_state, ssm_state\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        device = self.out_proj.weight.device\n        conv_dtype = self.conv1d.weight.dtype if dtype is None else dtype\n        conv_state = torch.zeros(\n            batch_size, self.d_model * self.expand, self.d_conv, device=device, dtype=conv_dtype\n        )\n        ssm_dtype = self.dt_proj.weight.dtype if dtype is None else dtype\n        # ssm_dtype = torch.float32\n        ssm_state = torch.zeros(\n            batch_size, self.d_model * self.expand, self.d_state, device=device, dtype=ssm_dtype\n        )\n        return conv_state, ssm_state\n\n    def _get_states_from_cache(self, inference_params, batch_size, initialize_states=False):\n        assert self.layer_idx is not None\n        if self.layer_idx not in inference_params.key_value_memory_dict:\n            batch_shape = (batch_size,)\n            conv_state = torch.zeros(\n                batch_size,\n                self.d_model * self.expand,\n                self.d_conv,\n                device=self.conv1d.weight.device,\n                dtype=self.conv1d.weight.dtype,\n            )\n            ssm_state = torch.zeros(\n                batch_size,\n                self.d_model * self.expand,\n                self.d_state,\n                device=self.dt_proj.weight.device,\n                dtype=self.dt_proj.weight.dtype,\n                # dtype=torch.float32,\n            )\n            inference_params.key_value_memory_dict[self.layer_idx] = (conv_state, ssm_state)\n        else:\n            conv_state, ssm_state = inference_params.key_value_memory_dict[self.layer_idx]\n            # TODO: What if batch size changes between generation, and we reuse the same states?\n            if initialize_states:\n                conv_state.zero_()\n                ssm_state.zero_()\n        return conv_state, ssm_state\n\n\nclass Block(nn.Module):\n    def __init__(\n        self, dim, mixer_cls, norm_cls=nn.LayerNorm, fused_add_norm=False, residual_in_fp32=False\n    ):\n        \"\"\"\n        Simple block wrapping a mixer class with LayerNorm/RMSNorm and residual connection\"\n\n        This Block has a slightly different structure compared to a regular\n        prenorm Transformer block.\n        The standard block is: LN -> MHA/MLP -> Add.\n        [Ref: https://arxiv.org/abs/2002.04745]\n        Here we have: Add -> LN -> Mixer, returning both\n        the hidden_states (output of the mixer) and the residual.\n        This is purely for performance reasons, as we can fuse add and LayerNorm.\n        The residual needs to be provided (except for the very first block).\n        \"\"\"\n        super().__init__()\n        self.residual_in_fp32 = residual_in_fp32\n        self.fused_add_norm = fused_add_norm\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n        if self.fused_add_norm:\n            assert RMSNorm is not None, \"RMSNorm import fails\"\n            assert isinstance(\n                self.norm, (nn.LayerNorm, RMSNorm)\n            ), \"Only LayerNorm and RMSNorm are supported for fused_add_norm\"\n\n    def forward(\n        self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None\n    ):\n        r\"\"\"Pass the input through the encoder layer.\n\n        Args:\n            hidden_states: the sequence to the encoder layer (required).\n            residual: hidden_states = Mixer(LN(residual))\n        \"\"\"\n        if not self.fused_add_norm:\n            residual = (hidden_states + residual) if residual is not None else hidden_states\n            hidden_states = self.norm(residual.to(dtype=self.norm.weight.dtype))\n            if self.residual_in_fp32:\n                residual = residual.to(torch.float32)\n        else:\n            fused_add_norm_fn = rms_norm_fn if isinstance(self.norm, RMSNorm) else layer_norm_fn\n            hidden_states, residual = fused_add_norm_fn(\n                hidden_states,\n                self.norm.weight,\n                self.norm.bias,\n                residual=residual,\n                prenorm=True,\n                residual_in_fp32=self.residual_in_fp32,\n                eps=self.norm.eps,\n            )\n        hidden_states = self.mixer(hidden_states, inference_params=inference_params)\n        return hidden_states, residual\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        return self.mixer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype, **kwargs)\n\u5206\u6790\u5728visionMamba\u7684forward_features\u4e2d\uff0cmamba\u7684\u4f5c\u7528\u662f\u4ec0\u4e48\uff1f"},{"assistant_initiates":false,"asst_text":"\u5728 `VisionMamba` \u7c7b\u7684 `forward_features` \u65b9\u6cd5\u4e2d\u6d89\u53ca\u5230 `Mamba` \u7684\u90e8\u5206\u4e3b\u8981\u662f\u5728\u5faa\u73af\u4e2d\u5bf9\u5404\u4e2a\u5c42\u8fdb\u884c\u8fed\u4ee3\u5904\u7406\u65f6\u8c03\u7528\u7684\u3002\u4ee5\u4e0b\u662f\u4e0e `Mamba` \u7c7b\u76f8\u5173\u7684\u90e8\u5206\u4ee3\u7801\uff1a\n\n```python\n# \u521d\u59cb\u5316 residual \u548c hidden_states\nresidual = None\nhidden_states = x\n\n# \u904d\u5386\u6a21\u578b\u7684\u5c42\nif not self.if_bidirectional:  # \u82e5\u6a21\u578b\u4e0d\u662f\u53cc\u5411\u7684\n    for layer in self.layers:\n        if if_flip_img_sequences and self.if_rope:\n            hidden_states = hidden_states.flip([1])\n            if residual is not None:\n                residual = residual.flip([1])\n\n        # \u4f7f\u7528 rotary embedding (\u82e5\u542f\u7528) \u5904\u7406\u7279\u5f81\n        if self.if_rope:\n            hidden_states = self.rope(hidden_states)\n            if residual is not None and self.if_rope_residual:\n                residual = self.rope(residual)\n\n        if if_flip_img_sequences and self.if_rope:\n            hidden_states = hidden_states.flip([1])\n            if residual is not None:\n                residual = residual.flip([1])\n\n        # \u8c03\u7528 Mamba \u7c7b\u5728\u5f53\u524d\u5c42\u7684\u5b9e\u4f8b\n        hidden_states, residual = layer(\n            hidden_states, residual, inference_params=inference_params\n        )\n```\n\n\u5728\u5faa\u73af\u4e2d\uff0c\u6bcf\u4e2a `layer` \u5b9e\u9645\u4e0a\u662f `Block` \u7684\u5b9e\u4f8b\uff0c\u800c `Block` \u7c7b\u4e2d\u5305\u542b\u4e86\u4e00\u4e2a `Mamba` \u7684\u5b9e\u4f8b\u3002\u5728 `layer` \u8c03\u7528\u65f6\uff08\u8fd9\u4e5f\u662f\u6bcf\u4e2a `Block` \u5b9e\u4f8b\u7684\u8c03\u7528\uff09\uff0c\u5b9e\u9645\u4e0a\u8c03\u7528\u4e86 `Block` \u7c7b\u7684 `forward` \u65b9\u6cd5\uff1a\n\n```python\nhidden_states = self.mixer(hidden_states, inference_params=inference_params)\n```\n\n\u5728\u8fd9\u91cc\uff0c`self.mixer` \u662f `Mamba` \u7c7b\u7684\u5b9e\u4f8b\u3002\u56e0\u6b64\uff0c\u5728\u5faa\u73af\u4e2d\u8c03\u7528\u6bcf\u4e2a `layer` \u65f6\uff0c\u5c31\u7b49\u4e8e\u8c03\u7528\u4e86 `Mamba` \u7684 `forward` \u65b9\u6cd5\u3002\u8fd9\u4e00\u8c03\u7528\u8d1f\u8d23\u5904\u7406\u8f93\u5165\u7684\u7279\u5f81\uff0c\u6267\u884c\u72b6\u6001\u7a7a\u95f4\u6a21\u578b\u548c\u4e00\u7cfb\u5217\u6295\u5f71\u3001\u5377\u79ef\u64cd\u4f5c\uff0c\u7136\u540e\u8fd4\u56de\u5904\u7406\u540e\u7684\u7279\u5f81\u3002\u901a\u8fc7\u8fd9\u4e9b\u5904\u7406\uff0c`Mamba` \u53ef\u4ee5\u5728\u7279\u5f81\u7684\u7ef4\u5ea6\u3001\u4e0a\u4e0b\u6587\u6df7\u5408\u548c\u5e8f\u5217\u5efa\u6a21\u4e0a\u53d1\u6325\u5173\u952e\u4f5c\u7528\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":1446,"n_turns_in_conversation":22,"n_words":131,"source_row_offset":38501,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"700d098128da95069ee0862d","turn_index":2,"user_text":"\u5c06forward_features\u4e2d\u6d89\u53ca\u5230mamba\u7684\u90e8\u5206\u5217\u51fa\u6765\u3002"},{"assistant_initiates":false,"asst_text":"\u8981\u5c06 `LinearAttention` \u96c6\u6210\u5230\u7ed9\u5b9a\u7684 `VisionMamba` \u7f51\u7edc\u4e2d\uff0c\u60a8\u9700\u8981\u4e00\u4e2a\u6e05\u6670\u7684\u76ee\u6807\u548c\u9002\u5f53\u7684\u4f4d\u7f6e\u3002\u901a\u5e38\uff0c\u6211\u4eec\u5e0c\u671b\u5728\u7279\u5b9a\u4efb\u52a1\u6216\u6a21\u578b\u9700\u6c42\u7684\u80cc\u666f\u4e0b\u6574\u5408\u6ce8\u610f\u529b\u673a\u5236\uff0c\u8fd9\u53ef\u80fd\u4f1a\u5728\u5757\u7ea7\u522b\uff08block-level\uff09\u6216\u7f51\u7edc\u7ea7\u522b\uff08network-level\uff09\u8fdb\u884c\u5b9e\u73b0\u3002\n\n\u8003\u8651\u5230 `VisionMamba` \u7684\u73b0\u6709\u67b6\u6784\u8bbe\u8ba1\uff0c\u60a8\u53ef\u80fd\u60f3\u8981\u5c06 `LinearAttention` \u66ff\u4ee3\u6216\u6dfb\u52a0\u5230 Transformer \u4e2d\u7684\u67d0\u4e9b\u4f4d\u7f6e\u3002\u4ee5\u4e0b\u662f\u4e00\u4e9b\u6f5c\u5728\u7684\u96c6\u6210\u70b9\uff1a\n\n1. **\u66ff\u6362 `Mamba` \u6a21\u5757**\uff1a\n   - \u5982\u679c `LinearAttention` \u88ab\u8ba4\u4e3a\u662f\u6a21\u578b\u7684\u4e00\u79cd\u6838\u5fc3\u6ce8\u610f\u529b\u673a\u5236\uff0c\u60a8\u53ef\u4ee5\u5728 `create_block` \u51fd\u6570\u4e2d\u521b\u5efa\u533a\u5757\uff08block\uff09\u7684\u8fc7\u7a0b\u4e2d\uff0c\u7528 `LinearAttention` \u66ff\u4ee3 `Mamba`\u3002\n   - \u60a8\u9700\u8981\u786e\u4fdd `LinearAttention` \u548c `Mamba` \u5177\u6709\u517c\u5bb9\u7684\u8f93\u5165\u8f93\u51fa\u89c4\u6a21\u3002\n\n2. **\u4f5c\u4e3a `Block` \u7684\u989d\u5916\u7ec4\u4ef6**\uff1a\n   - \u4f5c\u4e3a\u4e00\u79cd\u66f4\u590d\u6742\u7684\u96c6\u6210\u65b9\u6cd5\uff0c\u60a8\u4e5f\u53ef\u4ee5\u8bbe\u8ba1 `Block` \u6765\u540c\u65f6\u4f7f\u7528 `Mamba` \u548c `LinearAttention`\u3002\u8fd9\u6837\u60a8\u53ef\u4ee5\u901a\u8fc7\u5c06\u5b83\u4eec\u7684\u8f93\u51fa\u7ec4\u5408\u5728\u4e00\u8d77\u6216\u5e94\u7528\u67d0\u79cd\u901a\u8def\u9009\u62e9\u673a\u5236\u6765\u5229\u7528\u8fd9\u4e24\u8005\u7684\u4f18\u52bf\u3002\n   - \u5728\u8fd9\u79cd\u65b9\u6cd5\u4e2d\uff0c\u60a8\u9700\u8981\u8c03\u6574 `Block` \u7684 `forward` \u65b9\u6cd5\u903b\u8f91\uff0c\u4ee5\u4fbf\u5904\u7406\u6765\u81ea `LinearAttention` \u548c `Mamba` \u7684\u4e0d\u540c\u4fe1\u53f7\u3002\n\n3. **\u72ec\u7acb\u5c42**\uff1a\n   - \u5728 `VisionMamba` \u7f51\u7edc\u4e2d\uff0c\u53ef\u4ee5\u8003\u8651\u5728 `PatchEmbed` \u548c `Block` \u4e4b\u95f4\u6216\u8005\u662f\u67d0\u4e9b\u5176\u4ed6\u4f4d\u7f6e\u5355\u72ec\u63d2\u5165\u4e00\u4e2a `LinearAttention` \u5c42\u3002\n   - \u8fd9\u79cd\u65b9\u6cd5\u9002\u7528\u4e8e\u5728\u589e\u5f3a\u957f\u8ddd\u79bb\u7279\u5f81\u4ea4\u4e92\u65f6\u5e0c\u671b\u5f15\u5165\u989d\u5916\u7684\u6ce8\u610f\u529b\u6a21\u5757\uff0c\u800c\u4e0d\u5e72\u6270 `Mamba` \u7684\u6838\u5fc3\u529f\u80fd\u3002\n\n\u4ee5\u4e0b\u662f\u53ef\u80fd\u7684\u4ee3\u7801\u96c6\u6210\u793a\u4f8b\uff0c\u5047\u8bbe\u6211\u4eec\u5c06 `LinearAttention` \u89c6\u4e3a\u66ff\u4ee3\u4e00\u79cd\u6216\u6dfb\u52a0\u5230\u67d0\u79cd\u5b58\u5728\u7684\u673a\u5236\u4e2d\uff0c\u5982\uff1a\n\n```python\nclass Block(nn.Module):\n    def __init__(self, dim, mixer_cls, attention_cls=None, norm_cls=nn.LayerNorm, ...):\n        super().__init__()\n        self.attention = attention_cls(dim, input_resolution=(14, 14), num_heads=8) if attention_cls else None\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n        ...\n\n    def forward(self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None):\n        if self.attention:\n            hidden_states = self.attention(hidden_states)\n        \n        # Existing logic\n        hidden_states, residual = self.mixer(hidden_states, residual, inference_params=inference_params)\n        return hidden_states, residual\n\ndef create_block(...):\n    ...\n    # Create a block with both Mamba and LinearAttention\n    block = Block(\n        d_model,\n        mixer_cls=mixer_cls,\n        attention_cls=LinearAttention,  # Add this\n        norm_cls=norm_cls,\n        ...\n    )\n    ...\n```\n\n\u8bf7\u6839\u636e\u9700\u8981\u8c03\u6574\u548c\u6269\u5c55\u4ee3\u7801\u793a\u4f8b\uff0c\u4ee5\u786e\u4fdd\u8f93\u5165/\u8f93\u51fa\u7ef4\u5ea6\u7684\u517c\u5bb9\u6027\uff0c\u5e76\u8fdb\u884c\u9002\u5f53\u7684\u8bad\u7ec3\u548c\u6d4b\u8bd5\u4ee5\u9a8c\u8bc1\u96c6\u6210\u7684\u6548\u679c\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":3907,"n_turns_in_conversation":22,"n_words":412,"source_row_offset":38502,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"25d2bb9847b868b3b0d09806","turn_index":3,"user_text":"\u73b0\u5728\u57fa\u4e8e\u4e0a\u6587\uff0c\u6211\u4eec\u671f\u671b\u5c06class LinearAttention(nn.Module):\n    r\"\"\" Linear Attention with LePE and RoPE.\n\n    Args:\n        dim (int): Number of input channels.\n        num_heads (int): Number of attention heads.\n        qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, num_heads, qkv_bias=True, **kwargs):\n\n        super().__init__()\n        self.dim = dim\n        self.input_resolution = input_resolution\n        self.num_heads = num_heads\n        self.qk = nn.Linear(dim, dim * 2, bias=qkv_bias)\n        self.elu = nn.ELU()\n        self.lepe = nn.Conv2d(dim, dim, 3, padding=1, groups=dim)\n        self.rope = RoPE(shape=(input_resolution[0], input_resolution[1], dim))\n\n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: input features with shape of (B, N, C)\n        \"\"\"\n        b, n, c = x.shape\n        h = int(n ** 0.5)\n        w = int(n ** 0.5)\n        num_heads = self.num_heads\n        head_dim = c // num_heads\n\n        qk = self.qk(x).reshape(b, n, 2, c).permute(2, 0, 1, 3)\n        q, k, v = qk[0], qk[1], x\n        # q, k, v: b, n, c\n\n        q = self.elu(q) + 1.0\n        k = self.elu(k) + 1.0\n        q_rope = self.rope(q.reshape(b, h, w, c)).reshape(b, n, num_heads, head_dim).permute(0, 2, 1, 3)\n        k_rope = self.rope(k.reshape(b, h, w, c)).reshape(b, n, num_heads, head_dim).permute(0, 2, 1, 3)\n        q = q.reshape(b, n, num_heads, head_dim).permute(0, 2, 1, 3)\n        k = k.reshape(b, n, num_heads, head_dim).permute(0, 2, 1, 3)\n        v = v.reshape(b, n, num_heads, head_dim).permute(0, 2, 1, 3)\n\n        z = 1 / (q @ k.mean(dim=-2, keepdim=True).transpose(-2, -1) + 1e-6)\n        kv = (k_rope.transpose(-2, -1) * (n ** -0.5)) @ (v * (n ** -0.5))\n        x = q_rope @ kv * z\n\n        x = x.transpose(1, 2).reshape(b, n, c)\n        v = v.transpose(1, 2).reshape(b, h, w, c).permute(0, 3, 1, 2)\n        x = x + self.lepe(v).permute(0, 2, 3, 1).reshape(b, n, c)\n\n        return x\n\u96c6\u6210\u8fdb\u53bb\uff0c\u7ed9\u51fa\u96c6\u6210\u4f4d\u7f6e\u3002"},{"assistant_initiates":false,"asst_text":"\u8981\u7ec4\u5408 `Mamba` \u548c `LinearAttention`\uff0c\u540c\u65f6\u5229\u7528\u5b83\u4eec\u5404\u81ea\u7684\u7279\u6027\uff0c\u53ef\u4ee5\u5728 `Block` \u4e2d\u6784\u5efa\u4e00\u4e2a\u65b0\u7684\u7ed3\u6784\uff0c\u4f7f\u4e24\u8005\u80fd\u591f\u534f\u4f5c\u5de5\u4f5c\u3002\u8fd9\u79cd\u7ec4\u5408\u53ef\u4ee5\u4e3a\u6a21\u578b\u63d0\u4f9b\u66f4\u591a\u6837\u5316\u7684\u4fe1\u606f\u62bd\u8c61\u548c\u66f4\u5f3a\u7684\u8868\u73b0\u80fd\u529b\u3002\u4ee5\u4e0b\u662f\u4e00\u4e2a\u53ef\u80fd\u7684\u7ec4\u5408\u5b9e\u73b0\u65b9\u6848\uff1a\n\n### \u8bbe\u8ba1\u601d\u8def\n\n1. **\u6dfb\u52a0 `LinearAttention` \u5230 `Block`**\uff1a\n   - \u901a\u8fc7\u5728 `Block` \u4e2d\u5f15\u5165 `LinearAttention` \u5c42\uff0c\u60a8\u53ef\u4ee5\u5728 `Mamba` \u7684\u5904\u7406\u903b\u8f91\u4e4b\u524d\u6216\u4e4b\u540e\u8fdb\u884c\u7279\u5f81\u4ea4\u4e92\u3002\n   - \u8fd9\u5c06\u4f7f\u5f97\u6ce8\u610f\u529b\u673a\u5236\u4e0e\u72b6\u6001\u7a7a\u95f4\u6a21\u578b\u7684\u80fd\u529b\u76f8\u7ed3\u5408\u3002\n\n2. **\u7ec4\u5408\u4fe1\u606f**\uff1a\n   - \u53ef\u4ee5\u901a\u8fc7\u7b80\u5355\u7684\u52a0\u6cd5\u3001\u62fc\u63a5\u6216\u52a0\u6743\u6c42\u548c\u7b49\u7b56\u7565\u5c06 `LinearAttention` \u548c `Mamba` \u7684\u8f93\u51fa\u7ec4\u5408\u6210\u4e00\u4e2a\u7edf\u4e00\u7684\u4fe1\u606f\u8868\u793a\u3002\n   - \u4f7f\u7528\u9002\u5f53\u7684\u5f52\u4e00\u5316\u6216\u95e8\u63a7\u673a\u5236\u6765\u8c03\u8282\u8fd9\u4e24\u8005\u7684\u5f71\u54cd\u3002\n\n### \u4ee3\u7801\u793a\u4f8b\n\n\u7ed3\u5408\u4e24\u79cd\u673a\u5236\u540e\uff0c`Block` \u7c7b\u53ef\u80fd\u4f1a\u7c7b\u4f3c\u4e8e\u4ee5\u4e0b\u4ee3\u7801\uff1a\n\n```python\nclass Block(nn.Module):\n    def __init__(self, dim, mixer_cls, attention_cls, norm_cls=nn.LayerNorm, method='add', **kwargs):\n        \"\"\"\n        method: \u6307\u5b9a\u7ec4\u5408\u4e24\u8005\u8f93\u51fa\u7684\u65b9\u5f0f\uff0c\u4f8b\u5982 'add', 'concatenate'\n        \"\"\"\n        super().__init__()\n        self.attention = attention_cls(dim, input_resolution=(14, 14), num_heads=8, **kwargs)\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n        self.method = method\n\n    def forward(self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None):\n        attention_output = self.attention(hidden_states)\n        mixer_output, residual = self.mixer(hidden_states, residual, inference_params=inference_params)\n\n        if self.method == 'add':\n            combined_output = attention_output + mixer_output\n        elif self.method == 'concatenate':\n            combined_output = torch.cat((attention_output, mixer_output), dim=-1)\n            # If using concatenation, you may need an additional projection layer\n            combined_output = self.norm(combined_output)\n        \n        return combined_output, residual\n\ndef create_block(...):\n    ...\n    # Instantiate Block with both Mamba and LinearAttention\n    block = Block(\n        dim=d_model,\n        mixer_cls=mixer_cls,\n        attention_cls=LinearAttention,\n        norm_cls=norm_cls,\n        method='add',  # \u6216 'concatenate'\n        ...\n    )\n    ...\n```\n\n### \u6ce8\u610f\u4e8b\u9879\n\n- **\u517c\u5bb9\u6027**\uff1a\u8bf7\u786e\u4fdd `attention_output` \u548c `mixer_output` \u5728\u76f8\u52a0\u6216\u62fc\u63a5\u65f6\u5177\u6709\u76f8\u540c\u6216\u517c\u5bb9\u7684\u7279\u5f81\u7ef4\u5ea6\uff0c\u5c24\u5176\u662f\u5f53\u5408\u5e76\u65b9\u6cd5\u662f `\"concatenate\"` \u65f6\u3002\n- **\u6027\u80fd\u8c03\u4f18**\uff1a\u53ef\u4ee5\u901a\u8fc7\u6dfb\u52a0\u53c2\u6570\u8c03\u8282\u673a\u5236\uff08\u4f8b\u5982 `gate` \u673a\u5236\uff09\u6765\u7cbe\u7ec6\u63a7\u5236 `LinearAttention` \u548c `Mamba` \u7684\u5f71\u54cd\u7a0b\u5ea6\u3002\n- **\u5b9e\u9a8c\u9a8c\u8bc1**\uff1a\u8bb0\u5f97\u5728\u8bad\u7ec3\u548c\u6d4b\u8bd5\u9636\u6bb5\u5c0f\u5fc3\u9a8c\u8bc1\u548c\u8c03\u4f18\u7ec4\u5408\u6a21\u578b\uff0c\u4ee5\u786e\u4fdd\u8be5\u7ec4\u5408\u5728\u7279\u5b9a\u4efb\u52a1\u4e2d\u5e26\u6765\u6027\u80fd\u63d0\u5347\u3002\n\n\u901a\u8fc7\u8fd9\u79cd\u65b9\u5f0f\uff0c\u60a8\u65e2\u5229\u7528\u4e86 `LinearAttention` \u7684\u5f3a\u5927\u5c40\u90e8\u7279\u5f81\u63d0\u53d6\u80fd\u529b\uff0c\u53c8\u7ed3\u5408\u4e86 `Mamba` \u5728\u5904\u7406\u957f\u5e8f\u5217\u6216\u590d\u6742\u65f6\u5e8f\u4fe1\u53f7\u4e2d\u7684\u4f18\u52bf\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":2191,"n_turns_in_conversation":22,"n_words":176,"source_row_offset":38503,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"4e25993cdf89fadebd22c128","turn_index":4,"user_text":"\u6211\u4eec\u671f\u671b\u80fd\u591f\u7ec4\u5408mamba\u548cLinearAttention\uff0c\u800c\u4e0d\u662f\u53d6\u4ee3mamba"},{"assistant_initiates":false,"asst_text":"\u5f53\u7136\u53ef\u4ee5\u5c06 `Mamba` \u548c `LinearAttention` \u4e32\u8054\uff0c\u4ee5\u4fbf\u5728\u7279\u5f81\u5904\u7406\u4e2d\u4f9d\u6b21\u5e94\u7528\u8fd9\u4e24\u79cd\u4e0d\u540c\u7684\u64cd\u4f5c\u3002\u8fd9\u79cd\u4e32\u8054\u65b9\u5f0f\u5141\u8bb8\u6a21\u578b\u9996\u5148\u5229\u7528\u4e00\u79cd\u65b9\u6cd5\u5bf9\u8f93\u5165\u8fdb\u884c\u521d\u6b65\u5904\u7406\u540e\uff0c\u518d\u901a\u8fc7\u53e6\u4e00\u79cd\u65b9\u6cd5\u8fdb\u884c\u8fdb\u4e00\u6b65\u62bd\u8c61\u548c\u63d0\u53d6\u3002\n\n\u4e0b\u9762\u662f\u5982\u4f55\u5728 `Block` \u7c7b\u4e2d\u5b9e\u73b0\u4e32\u8054\u64cd\u4f5c\u7684\u793a\u4f8b\uff1a\n\n### \u7f16\u7801\u793a\u4f8b\n\n\u5728\u8fd9\u4e2a\u793a\u4f8b\u4e2d\uff0c\u6211\u4eec\u5047\u8bbe\u5148\u8fdb\u884c `LinearAttention` \u5904\u7406\uff0c\u7136\u540e\u518d\u8fdb\u884c `Mamba` \u5904\u7406\uff1a\n\n```python\nclass Block(nn.Module):\n    def __init__(self, dim, mixer_cls, attention_cls, norm_cls=nn.LayerNorm, **kwargs):\n        super().__init__()\n        # \u521d\u59cb\u5316 LinearAttention \u548c Mamba \u4e24\u4e2a\u5bf9\u8c61\n        self.attention = attention_cls(dim, input_resolution=(14, 14), num_heads=8, **kwargs)\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n\n    def forward(self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None):\n        # \u5c06\u8f93\u5165\u7279\u5f81\u7ecf\u8fc7 LinearAttention \u5904\u7406\n        attention_output = self.attention(hidden_states)\n        \n        # \u4f7f\u7528\u5904\u7406\u540e\u7684\u7279\u5f81\u7ecf\u8fc7 Mamba\n        mixer_output, residual = self.mixer(attention_output, residual, inference_params=inference_params)\n        \n        # \u5c06\u6700\u7ec8\u8f93\u51fa\u8fdb\u884c\u5f52\u4e00\u5316\u6216\u5176\u4ed6\u540e\u5904\u7406\n        mixed_output = self.norm(mixer_output)\n        \n        return mixed_output, residual\n\ndef create_block(...):\n    ...\n    # \u521b\u5efa Block \u5bf9\u8c61\uff0c\u5305\u542b\u4e32\u8054\u7684 LinearAttention \u548c Mamba\n    block = Block(\n        dim=d_model,\n        mixer_cls=mixer_cls,\n        attention_cls=LinearAttention,\n        norm_cls=norm_cls,\n        ...\n    )\n    ...\n```\n\n### \u903b\u8f91\u8bf4\u660e\n\n1. **\u7ebf\u6027\u6ce8\u610f\u529b\uff08LinearAttention\uff09**\uff1a\n   - \u9996\u5148\u5bf9\u8f93\u5165\u7279\u5f81\u8fdb\u884c\u7ebf\u6027\u81ea\u6ce8\u610f\u529b\u64cd\u4f5c\u3002\u8fd9\u79cd\u65b9\u6cd5\u9002\u7528\u4e8e\u4ece\u8f93\u5165\u4e2d\u63d0\u53d6\u7cbe\u7ec6\u7684\u5c40\u90e8\u5173\u7cfb\u3002\n\n2. **\u72b6\u6001\u7a7a\u95f4\u6a21\u578b\uff08Mamba\uff09**\uff1a\n   - \u5728\u7ecf\u5386\u81ea\u6ce8\u610f\u529b\u5904\u7406\u540e\uff0c\u7ed3\u679c\u88ab\u4f20\u9012\u5230 `Mamba`\uff0c\u5229\u7528\u5176\u5728\u957f\u5e8f\u5217\u5efa\u6a21\u548c\u590d\u6742\u65f6\u5e8f\u5904\u7406\u4e2d\u7684\u4f18\u52bf\u3002\n\n3. **\u5f52\u4e00\u5316**\uff1a\n   - \u53ef\u4ee5\u9009\u62e9\u5728\u6700\u540e\u52a0\u4e0a\u5f52\u4e00\u5316\uff08\u6216\u5176\u4ed6\u540e\u5904\u7406\uff09\uff0c\u786e\u4fdd\u6570\u636e\u8303\u56f4\u7684\u7a33\u5b9a\u6027\u548c\u63d0\u9ad8\u8bad\u7ec3\u7684\u6536\u655b\u6027\u3002\n\n### \u4f7f\u7528\u6ce8\u610f\n\n- **\u987a\u5e8f**\uff1a\u6ce8\u610f `Mamba` \u548c `LinearAttention` \u7684\u987a\u5e8f\u53ef\u4ee5\u6839\u636e\u4efb\u52a1\u9700\u6c42\u8c03\u6574\u3002\n- **\u6027\u80fd**\uff1a\u901a\u8fc7\u8bad\u7ec3\u548c\u9a8c\u8bc1\u786e\u4fdd\u8fd9\u79cd\u4e32\u8054\u5e26\u6765\u7684\u6027\u80fd\u63d0\u5347\uff0c\u5e76\u6839\u636e\u5b9e\u9a8c\u7ed3\u679c\u5bf9\u6a21\u5757\u7684\u987a\u5e8f\u6216\u8d85\u53c2\u6570\u8fdb\u884c\u8c03\u4f18\u3002\n- **\u8f93\u5165\u8f93\u51fa\u7ef4\u5ea6**\uff1a\u786e\u4fdd\u6bcf\u6b65\u5904\u7406\u7684\u8f93\u5165\u8f93\u51fa\u7ef4\u5ea6\u662f\u5339\u914d\u7684\u3002\u5fc5\u8981\u65f6\uff0c\u5728\u4e0d\u540c\u6a21\u5757\u4e4b\u95f4\u52a0\u5165\u989d\u5916\u7684\u6295\u5f71\u5c42\u4ee5\u8c03\u6574\u7279\u5f81\u7ef4\u5ea6\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":1687,"n_turns_in_conversation":22,"n_words":127,"source_row_offset":38504,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"2bcc5495fffca559fe99abce","turn_index":5,"user_text":"\u90a3\u4e48\u80fd\u5426\u4e32\u8054\u4e24\u79cd\u64cd\u4f5c\uff1f"},{"assistant_initiates":false,"asst_text":"\u8981\u8bbe\u7f6e `LinearAttention` \u7684\u53c2\u6570\uff0c\u60a8\u9700\u8981\u786e\u4fdd\u5728 `Block` \u7c7b\u548c `create_block` \u51fd\u6570\u4e2d\u6b63\u786e\u5730\u4f20\u9012\u8fd9\u4e9b\u53c2\u6570\u3002\u8fd9\u91cc\u662f\u5982\u4f55\u7ec4\u7ec7\u548c\u8bbe\u7f6e\u53c2\u6570\u7684\u4e00\u4e9b\u5efa\u8bae\uff1a\n\n### \u53c2\u6570\u8bbe\u7f6e\n\n1. **\u786e\u5b9a `LinearAttention` \u7684\u8f93\u5165\u53c2\u6570**\uff1a\n   - `dim`: \u8f93\u5165\u901a\u9053\u7684\u6570\u91cf\u901a\u5e38\u5bf9\u5e94\u4e8e\u6a21\u578b\u7684\u5d4c\u5165\u7ef4\u5ea6\u3002\n   - `input_resolution`: \u8fd9\u662f\u4e00\u4e2a\u4e8c\u5143\u7ec4\uff0c\u8868\u793a\u8f93\u5165\u7279\u5f81\u56fe\u7684\u957f\u548c\u5bbd\u3002\u8fd9\u53ef\u80fd\u9700\u8981\u4ece\u5916\u90e8\u4f20\u5165\uff0c\u6216\u6839\u636e\u6a21\u578b\u7279\u5b9a\u9700\u6c42\u8bbe\u7f6e\u3002\n   - `num_heads`: \u6ce8\u610f\u529b\u5934\u7684\u6570\u91cf\uff1b\u53ef\u4ee5\u662f\u56fa\u5b9a\u503c\u6216\u4f5c\u4e3a\u53c2\u6570\u4f20\u5165\u3002\n\n2. **\u5229\u7528 `**kwargs` \u6269\u5c55\u4f20\u5165\u7684\u5176\u4ed6\u53c2\u6570**\uff1a\n   - \u5982\u679c `LinearAttention` \u7c7b\u6709\u5176\u4ed6\u53ef\u9009\u6216\u989d\u5916\u7684\u53c2\u6570\uff0c\u53ef\u4ee5\u901a\u8fc7 `**kwargs` \u4f20\u4e34\u65f6\u53c2\u6570\u3002\n\n### \u793a\u4f8b\u4ee3\u7801\n\n\u4ee5\u4e0b\u662f\u5982\u4f55\u5728 `Block` \u7c7b\u4e2d\u4f20\u9012\u53c2\u6570\u5e76\u5728 `create_block` \u4e2d\u8c03\u7528\uff1a\n\n```python\nclass Block(nn.Module):\n    def __init__(self, dim, mixer_cls, attention_cls, input_resolution, num_heads=8, norm_cls=nn.LayerNorm, **kwargs):\n        super().__init__()\n        # \u521d\u59cb\u5316 LinearAttention \u548c Mamba\n        self.attention = attention_cls(dim, input_resolution=input_resolution, num_heads=num_heads, **kwargs)\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n\n    def forward(self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None):\n        # \u901a\u8fc7 LinearAttention \u5904\u7406\n        attention_output = self.attention(hidden_states)\n        \n        # \u7ecf\u8fc7 Mamba \u5904\u7406\n        mixer_output, residual = self.mixer(attention_output, residual, inference_params=inference_params)\n        \n        # \u6700\u540e\u7684\u5904\u7406\n        mixed_output = self.norm(mixer_output)\n        \n        return mixed_output, residual\n\ndef create_block(d_model, ssm_cfg=None, norm_epsilon=1e-5, drop_path=0., rms_norm=False, input_resolution=(14, 14), num_heads=8, **kwargs):\n    # \u8bbe\u7f6e\u5de5\u5382\u53c2\u6570\n    factory_kwargs = {\"device\": device, \"dtype\": dtype}\n    \n    # \u5b9a\u4e49 mixer (Mamba) \u7c7b\n    mixer_cls = partial(Mamba, **ssm_cfg, **factory_kwargs)\n    \n    # \u5b9a\u4e49 norm \u7c7b\n    norm_cls = partial(nn.LayerNorm if not rms_norm else RMSNorm, eps=norm_epsilon, **factory_kwargs)\n\n    # \u521b\u5efa Block \u5bf9\u8c61\n    block = Block(\n        dim=d_model,\n        mixer_cls=mixer_cls,\n        attention_cls=LinearAttention,\n        input_resolution=input_resolution,\n        num_heads=num_heads,\n        norm_cls=norm_cls,\n        **kwargs\n    )\n    return block\n```\n\n### \u6ce8\u610f\u4e8b\u9879\n\n- **\u53c2\u6570\u5316**\uff1a\u786e\u4fdd `input_resolution` \u548c `num_heads` \u662f\u4e0e\u4efb\u52a1\u76f8\u5173\u7684\u9002\u5f53\u503c\uff0c\u60a8\u53ef\u4ee5\u6839\u636e\u7279\u5b9a\u7684\u6a21\u578b\u548c\u6570\u636e\u96c6\u9700\u6c42\u8c03\u6574\u8fd9\u4e9b\u53c2\u6570\u3002\n- **\u6269\u5c55\u6027**\uff1a\u901a\u8fc7\u4f7f\u7528 `**kwargs`\uff0c\u53ef\u4ee5\u65b9\u4fbf\u5730\u5411 `LinearAttention` \u6216 `Block` \u7c7b\u4f20\u9012\u989d\u5916\u6216\u65b0\u7684\u53c2\u6570\uff0c\u800c\u65e0\u9700\u4fee\u6539\u73b0\u6709\u51fd\u6570\u7b7e\u540d\u3002\n- **\u8c03\u8bd5**\uff1a\u9a8c\u8bc1\u53c2\u6570\u7684\u6b63\u786e\u4f20\u9012\u548c\u8fd0\u884c\u4e3a\u4fdd\u8bc1\u6a21\u578b\u5728\u521d\u671f\u9636\u6bb5\u6b63\u786e\u8bbe\u7f6e\u7684\u5173\u952e\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":2296,"n_turns_in_conversation":22,"n_words":187,"source_row_offset":38505,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"9f218a436a8aab922b4d2a0a","turn_index":6,"user_text":"\u90a3\u4e48\u73b0\u5728\u9700\u8981\u8bbe\u7f6e\u7684\u662f\u4f20\u53c2\u7684\u95ee\u9898.self.attention = attention_cls(dim, input_resolution=(14, 14), num_heads=8, **kwargs)\u5982\u4f55\u8bbe\u7f6e\uff1f"},{"assistant_initiates":false,"asst_text":"\u4e3a\u4e86\u5c06 `Mamba` \u548c `LinearAttention` \u7684\u8f93\u51fa\u76f8\u52a0\uff0c\u60a8\u53ef\u4ee5\u5728 `Block` \u7c7b\u4e2d\u5b9a\u4e49\u4e00\u79cd\u7ec4\u5408\u65b9\u6cd5\uff0c\u5728\u5c06\u8f93\u5165\u4f20\u9012\u901a\u8fc7\u8fd9\u4e24\u4e2a\u6a21\u5757\u540e\uff0c\u5c06\u5b83\u4eec\u7684\u8f93\u51fa\u76f8\u52a0\u3002\u4ee5\u4e0b\u662f\u5b9e\u73b0\u8fd9\u4e00\u903b\u8f91\u7684\u793a\u4f8b\u4ee3\u7801\uff1a\n\n### \u7f16\u7801\u6b65\u9aa4\n\n1. **\u5728 `Block` \u4e2d\u5b9a\u4e49\u4e24\u4e2a\u6a21\u5757**\uff1a\u4e00\u4e2a\u7528\u4e8e `LinearAttention`\uff0c\u53e6\u4e00\u4e2a\u7528\u4e8e `Mamba`\u3002\n2. **\u5728 `forward` \u65b9\u6cd5\u4e2d\u5904\u7406\u8f93\u5165**\uff1a\u5c06\u8f93\u5165\u4f9d\u6b21\u4f20\u9012\u7ed9 `LinearAttention` \u548c `Mamba`\u3002\n3. **\u7ec4\u5408\u64cd\u4f5c**\uff1a\u5c06\u8fd9\u4e24\u8005\u7684\u8f93\u51fa\u6dfb\u52a0\u5728\u4e00\u8d77\uff0c\u5f62\u6210\u6700\u7ec8\u7684\u8f93\u51fa\u3002\n\n### \u793a\u4f8b\u4ee3\u7801\n\n\u4ee5\u4e0b\u662f\u5982\u4f55\u5b9e\u73b0\u8fd9\u4e00\u65b9\u6848\u7684\u793a\u4f8b\uff1a\n\n```python\nclass Block(nn.Module):\n    def __init__(self, dim, mixer_cls, attention_cls, input_resolution, num_heads=8, norm_cls=nn.LayerNorm, **kwargs):\n        super().__init__()\n        # \u521d\u59cb\u5316 LinearAttention \u548c Mamba \u6a21\u5757\n        self.attention = attention_cls(dim, input_resolution=input_resolution, num_heads=num_heads, **kwargs)\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n\n    def forward(self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None):\n        # \u7ecf\u8fc7 LinearAttention \u5904\u7406\n        attention_output = self.attention(hidden_states)\n        \n        # \u7ecf\u8fc7 Mamba \u5904\u7406\n        mamba_output, residual = self.mixer(hidden_states, residual, inference_params=inference_params)\n        \n        # \u5c06\u4e24\u8005\u7684\u8f93\u51fa\u76f8\u52a0\n        combined_output = attention_output + mamba_output\n        \n        # \u8fdb\u884c\u5f52\u4e00\u5316\n        normalized_output = self.norm(combined_output)\n        \n        return normalized_output, residual\n\ndef create_block(d_model, ssm_cfg=None, norm_epsilon=1e-5, drop_path=0., rms_norm=False, input_resolution=(14, 14), num_heads=8, **kwargs):\n    # \u8bbe\u7f6e\u5de5\u5382\u53c2\u6570\n    factory_kwargs = {\"device\": kwargs.get('device'), \"dtype\": kwargs.get('dtype')}\n    \n    # \u5b9a\u4e49 mixer (Mamba) \u7c7b\n    mixer_cls = partial(Mamba, **ssm_cfg, **factory_kwargs)\n    \n    # \u5b9a\u4e49 norm \u7c7b\n    norm_cls = partial(nn.LayerNorm if not rms_norm else RMSNorm, eps=norm_epsilon, **factory_kwargs)\n\n    # \u521b\u5efa Block \u5bf9\u8c61\n    block = Block(\n        dim=d_model,\n        mixer_cls=mixer_cls,\n        attention_cls=LinearAttention,\n        input_resolution=input_resolution,\n        num_heads=num_heads,\n        norm_cls=norm_cls,\n        **kwargs\n    )\n    return block\n```\n\n### \u91cd\u8981\u6ce8\u610f\u4e8b\u9879\n\n- **\u8f93\u5165\u7ef4\u5ea6\u5339\u914d**\uff1a\u786e\u4fdd\u8f93\u5165\u5230 `attention` \u548c `mamba` \u7684\u7279\u5f81\u7ef4\u5ea6\u662f\u5339\u914d\u7684\u3002\u8fd9\u79cd\u76f8\u52a0\u64cd\u4f5c\u8981\u6c42\u8fd9\u4e24\u8005\u8f93\u51fa\u7279\u5f81\u5177\u6709\u76f8\u540c\u7684\u5f62\u72b6\u3002\n  \n- **\u8c03\u8bd5\u4e0e\u9a8c\u8bc1**\uff1a\u5728\u6a21\u578b\u90e8\u7f72\u5230\u5b9e\u9a8c\u4e2d\u524d\uff0c\u4ed4\u7ec6\u9a8c\u8bc1\u7f51\u7edc\u7684\u521d\u59cb\u5316\u548c\u524d\u5411\u4f20\u64ad\uff0c\u4ee5\u786e\u8ba4\u7ec4\u5408\u540e\u7684\u6a21\u578b\u5728\u60a8\u7684\u8bad\u7ec3\u4efb\u52a1\u4e2d\u8868\u73b0\u826f\u597d\u3002\u53ef\u80fd\u9700\u8981\u5bf9\u8d85\u53c2\u6570\u8fdb\u884c\u8c03\u8bd5\u4ee5\u8fbe\u5230\u6700\u4f73\u6548\u679c\u3002\n\n\u901a\u8fc7\u8fd9\u79cd\u65b9\u6cd5\uff0c\u60a8\u53ef\u4ee5\u6709\u6548\u5730\u7ed3\u5408 `Mamba` \u548c `LinearAttention` \u7684\u4f18\u52bf\uff0c\u4e24\u8005\u7684\u5171\u540c\u4f5c\u7528\u6709\u671b\u63d0\u5347\u7279\u5f81\u63d0\u53d6\u548c\u5efa\u6a21\u80fd\u529b\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":2228,"n_turns_in_conversation":22,"n_words":175,"source_row_offset":38506,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"c35c3ce9944d0da3f57aaba4","turn_index":7,"user_text":"\u8fd9\u6837\u5427\uff0c\u628a\u4e8c\u8005add\u8d77\u6765"},{"assistant_initiates":false,"asst_text":"\u5728\u7ed9\u5b9a\u4ee3\u7801\u4e2d\uff0c`LinearAttention` \u5df2\u88ab\u96c6\u6210\u5230 `Block` \u7c7b\u4e2d\u3002\u6211\u5c06\u6ce8\u610f\u5230\u51e0\u4e2a\u53ef\u80fd\u7684\u95ee\u9898\u548c\u9700\u8981\u4fee\u6b63\u7684\u5730\u65b9\uff0c\u4ee5\u53ca\u9700\u8981\u6ce8\u610f\u7684\u65b9\u9762\uff1a\n\n### 1. **\u7f3a\u5931\u7684\u53c2\u6570\u4f20\u9012**\n\n\u5728 `Block` \u7c7b\u7684\u6784\u9020\u51fd\u6570\u4e2d\uff0c`LinearAttention` \u88ab\u521d\u59cb\u5316\u65f6\u4f7f\u7528\u4e86 `**kwargs`\uff1a\n\n```python\nself.attention = attention_cls(dim, input_resolution=input_resolution, num_heads=num_heads, **kwargs)\n```\n\n\u4e0d\u8fc7\uff0c\u8fd9\u4e2a `kwargs` \u5728 `Block` \u7684\u521d\u59cb\u5316\u4e2d\u4f3c\u4e4e\u6ca1\u6709\u88ab\u663e\u5f0f\u4f20\u9012\uff0c\u56e0\u6b64\u4f20\u5165\u7684\u53c2\u6570\u5fc5\u987b\u5728 `create_block` \u51fd\u6570\u6216 `Block` \u521d\u59cb\u5316\u65f6\u660e\u786e\u4f20\u9012\u3002\n\n### 2. **\u8f93\u5165/\u8f93\u51fa\u7ef4\u5ea6**\n\n\u786e\u4fdd `LinearAttention` \u548c `Mamba` \u7684\u8f93\u5165\u548c\u8f93\u51fa\u7ef4\u5ea6\u662f\u76f8\u5339\u914d\u7684\u3002\u5728\u524d\u5411\u4f20\u64ad\u4e2d\uff1a\n\n```python\nattention_out = self.attention(hidden_states)\nhidden_states = self.mixer(hidden_states, inference_params=inference_params)\n```\n\n\u8bf7\u786e\u4fdd `hidden_states` \u7684\u7ef4\u5ea6\u5bf9\u4e8e\u8fd9\u4e24\u4e2a\u6a21\u5757\u90fd\u662f\u76f8\u540c\u7684\uff0c\u4ee5\u4f7f\u5f97\u5b83\u4eec\u80fd\u591f\u987a\u5229\u5904\u7406\u3002\u5982\u679c `hidden_states` \u7684\u7ef4\u5ea6\u5728 `attention` \u6216 `mixer` \u5904\u7406\u540e\u53d1\u751f\u6539\u53d8\uff08\u901a\u5e38\u4e0d\u4f1a\uff0c\u56e0\u4e3a\u5b83\u4eec\u90fd\u4e3a\u7ebf\u6027\u6295\u5f71\uff09\uff0c\u53ef\u80fd\u9700\u8981\u989d\u5916\u7684\u7ebf\u6027\u5c42\u6765\u5339\u914d\u8fd9\u4e2a\u53d8\u5316\u3002\n\n### 3. **GLU \u96c6\u6210\u903b\u8f91**\n\n`ResGLUAdapter` \u53ef\u80fd\u9700\u8981\u660e\u786e\u7684\u521d\u59cb\u5316\uff0c\u5b83\u5728\u4ee3\u7801\u4e2d\u88ab\u6761\u4ef6\u6027\u5f15\u5165\uff1a\n\n```python\nif integrate_glu:\n    self.glu_adapter = ResGLUAdapter(model_args)\nelse:\n    self.glu_adapter = None \n```\n\n\u786e\u4fdd `ResGLUAdapter` \u7684\u96c6\u6210\u903b\u8f91\u6e05\u6670\u53ef\u9a8c\u8bc1\uff0c\u6309\u60a8\u6240\u9700\u7684\u6761\u4ef6\u81ea\u52a8\u521d\u59cb\u5316\u3002\u6b64\u5916\uff0c\u9700\u8981\u9a8c\u8bc1\u7684\u662f `ResGLUAdapter` \u7684\u8f93\u5165\u548c\u8f93\u51fa\u7ef4\u5ea6\u5408\u4e4e\u4e0a\u4e0b\u6587\u903b\u8f91\u3002\n\n### 4. **forward_features \u65b9\u6cd5**\n\n\u5728\u4f7f\u7528 `LinearAttention`\uff0c`Mamba` \u548c\u53ef\u80fd\u7684 `GLU` \u65f6\uff1a\n\n```python\n# In Block.forward\nattention_out = self.attention(hidden_states)\nhidden_states = self.mixer(hidden_states, inference_params=inference_params)\nif self.glu_adapter is not None:\n    hidden_states = self.glu_adapter(hidden_states)\nhidden_states = attention_out + hidden_states\n```\n\n\u786e\u4fdd `attention_out` \u548c `hidden_states` \u7684\u76f8\u52a0\u64cd\u4f5c\u5728\u5f62\u72b6\u548c\u6e20\u9053\u6570\u91cf\u4e0a\u662f\u76f8\u5bb9\u7684\u3002\u901a\u5e38\uff0c\u5728 Transformer \u67b6\u6784\u4e2d\u8fd9\u4f1a\u5f88\u81ea\u7136\u5730\u5f97\u5230\u6ee1\u8db3\uff0c\u7136\u800c\uff0c\u5fc5\u987b\u786e\u8ba4\u8fd9\u4e00\u70b9\u3002\n\n### 5. **\u53c2\u6570\u521d\u59cb\u5316\u548c\u6743\u91cd\u5171\u4eab**\n\n\u4efb\u4f55\u9700\u8981\u5171\u4eab\u7684\u53c2\u6570\u6216\u6743\u91cd\u9700\u8981\u786e\u4fdd\u5728\u591a\u4e2a `Block` \u6216\u6574\u4e2a `VisionMamba` \u6a21\u578b\u4e2d\u8ba9\u6743\u91cd\u4e4b\u95f4\u4fdd\u6301\u4e00\u81f4\u3002\n\n### \u63a8\u8350\u64cd\u4f5c\n\n- **\u6d4b\u8bd5**\uff1a\u901a\u8fc7\u5355\u5143\u6d4b\u8bd5\u786e\u4fdd\u6bcf\u4e2a\u5757\u4ee5\u53ca\u7ec4\u5408\u6a21\u5757\u529f\u80fd\u7684\u76f8\u4e92\u8f6c\u6362\u3002\n- **\u65e5\u5fd7\u8bb0\u5f55**\uff1a\u589e\u52a0\u8c03\u8bd5\u4fe1\u606f\u4ee5\u5e2e\u52a9\u786e\u8bca\u95ee\u9898\u3002\u53ef\u4ee5\u5728 forward \u51fd\u6570\u4e2d\u901a\u8fc7\u6253\u5370\u5f20\u91cf\u7684\u5c3a\u5bf8\u6765\u9a8c\u8bc1\u3002\n\n\u901a\u8fc7\u8fd9\u4e9b\u8c03\u8bd5\u548c\u6539\u8fdb\u63aa\u65bd\uff0c\u53ef\u4ee5\u5e2e\u52a9\u786e\u4fdd `LinearAttention` \u7684\u6b63\u786e\u96c6\u6210\uff0c\u540c\u65f6\u6709\u6548\u5b9e\u73b0\u5176\u7279\u6027\u63d0\u4f9b\u7684\u589e\u5f3a\u6548\u679c\u3002","behavior":"degenerate_repetition","conversation_hash":"fc57d6b868a013f3dd4e8b6abec3a8fa","conversation_redacted":false,"conversation_toxic":false,"country":"Singapore","empty_output":false,"evidence_codes":[],"label":"absent","label_present":false,"label_qualified":false,"label_source":"Final dense release","language":"Chinese","model":"gpt-4o-2024-08-06","n_chars":29083,"n_turns_in_conversation":22,"n_words":2124,"source_row_offset":38507,"source_shard":"turns-00042.parquet","timestamp":"Fri, 01 Nov 2024 09:34:30 GMT","turn_id":"949e375ee9149f0f0e4a45fd","turn_index":8,"user_text":"\u73b0\u5728\u5206\u6790import torch\nimport torch.nn as nn\nfrom functools import partial\nfrom torch import Tensor\nfrom typing import Optional\n\nfrom timm.models.vision_transformer import VisionTransformer, _cfg\nfrom timm.models.registry import register_model\nfrom timm.models.layers import trunc_normal_, lecun_normal_\n\nfrom timm.models.layers import DropPath, to_2tuple\nfrom timm.models.vision_transformer import _load_weights\n\nimport math\n\nfrom collections import namedtuple\n\nfrom mamba_ssm.modules.mamba_simple import Mamba\nfrom mamba_ssm.utils.generation import GenerationMixin\nfrom mamba_ssm.utils.hf import load_config_hf, load_state_dict_hf\n\nfrom lib.models.mamba_fetrack.rope import *\nimport random\nfrom lib.models.layers.head import build_box_head\nimport importlib\nimport lib.train.admin.settings as ws_settings\nfrom .utils import combine_tokens, recover_tokens\n\n# \nfrom .HinBlock import HinResBlock\nfrom .ResGLU import ResGLUAdapter, ModelArgs\nfrom .Attn import LinearAttention\ntry:\n    from mamba_ssm.ops.triton.layernorm import RMSNorm, layer_norm_fn, rms_norm_fn\nexcept ImportError:\n    RMSNorm, layer_norm_fn, rms_norm_fn = None, None, None\n\n__all__ = [\n    'vim_tiny_patch16_224', 'vim_small_patch16_224', 'vim_base_patch16_224',\n    'vim_tiny_patch16_384', 'vim_small_patch16_384', 'vim_base_patch16_384',\n]\n\n\nclass PatchEmbed(nn.Module):\n    \"\"\" 2D Image to Patch Embedding\n    \"\"\"\n    def __init__(self, img_size=224, patch_size=16, stride=16, in_chans=3, embed_dim=768, norm_layer=None, flatten=True):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.grid_size = ((img_size[0] - patch_size[0]) // stride + 1, (img_size[1] - patch_size[1]) // stride + 1)\n        self.num_patches = self.grid_size[0] * self.grid_size[1]\n        self.flatten = flatten\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=stride)\n        self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        x = self.proj(x)\n        if self.flatten:\n            x = x.flatten(2).transpose(1, 2)  # BCHW -> BNC\n        x = self.norm(x)\n        return x\n    \n\nclass Block(nn.Module):\n    def __init__(\n        self, dim, mixer_cls, attention_cls, input_resolution, num_heads=8, norm_cls=nn.LayerNorm, fused_add_norm=False, residual_in_fp32=False,drop_path=0., integrate_glu=False\n    ):\n        \"\"\"\n        Simple block wrapping a mixer class with LayerNorm/RMSNorm and residual connection\"\n\n        This Block has a slightly different structure compared to a regular\n        prenorm Transformer block.\n        The standard block is: LN -> MHA/MLP -> Add.\n        [Ref: https://arxiv.org/abs/2002.04745]\n        Here we have: Add -> LN -> Mixer, returning both\n        the hidden_states (output of the mixer) and the residual.\n        This is purely for performance reasons, as we can fuse add and LayerNorm.\n        The residual needs to be provided (except for the very first block).\n        \"\"\"\n        super().__init__()\n        self.residual_in_fp32 = residual_in_fp32\n        self.fused_add_norm = fused_add_norm\n        # \n        self.attention = attention_cls(dim, input_resolution=input_resolution, num_heads=num_heads, **kwargs)\n        self.mixer = mixer_cls(dim)\n        self.norm = norm_cls(dim)\n        #\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        if self.fused_add_norm:\n            assert RMSNorm is not None, \"RMSNorm import fails\"\n            assert isinstance(\n                self.norm, (nn.LayerNorm, RMSNorm)\n            ), \"Only LayerNorm and RMSNorm are supported for fused_add_norm\"\n        # \u5f15\u5165\u95e8\u63a7\u673a\u5236\n        model_args = ModelArgs()\n        # \u5982\u679c\u9700\u8981\u96c6\u6210 ResGLUAdapter\uff0c\u5219\u521d\u59cb\u5316\n        if integrate_glu:\n            model_args = ModelArgs()\n            self.glu_adapter = ResGLUAdapter(model_args)\n        else:\n            self.glu_adapter = None \n        # self.glu_adapter = ResGLUAdapter(model_args)\n\n    def forward(\n        self, hidden_states: Tensor, residual: Optional[Tensor] = None, inference_params=None\n    ):\n        r\"\"\"Pass the input through the encoder layer.\n\n        Args:\n            hidden_states: the sequence to the encoder layer (required).\n            residual: hidden_states = Mixer(LN(residual))\n        \"\"\"\n        if not self.fused_add_norm:\n            if residual is None:\n                residual = hidden_states\n            else:\n                residual = residual + self.drop_path(hidden_states)\n            \n            hidden_states = self.norm(residual.to(dtype=self.norm.weight.dtype))\n            if self.residual_in_fp32:\n                residual = residual.to(torch.float32)\n        else:\n            fused_add_norm_fn = rms_norm_fn if isinstance(self.norm, RMSNorm) else layer_norm_fn\n            if residual is None:\n                hidden_states, residual = fused_add_norm_fn(\n                    hidden_states,\n                    self.norm.weight,\n                    self.norm.bias,\n                    residual=residual,\n                    prenorm=True,\n                    residual_in_fp32=self.residual_in_fp32,\n                    eps=self.norm.eps,\n                )\n            else:\n                hidden_states, residual = fused_add_norm_fn(\n                    self.drop_path(hidden_states),\n                    self.norm.weight,\n                    self.norm.bias,\n                    residual=residual,\n                    prenorm=True,\n                    residual_in_fp32=self.residual_in_fp32,\n                    eps=self.norm.eps,\n                )  \n        attention_out = self.attention(hidden_states)  \n        hidden_states = self.mixer(hidden_states, inference_params=inference_params)\n        # \u589e\u52a0\u95e8\u63a7\u673a\u5236, \u6d88\u9664\u5197\u4f59\u548c\u4e0d\u76f8\u5173\u7684\u7279\u5f81\uff0c\u4fdd\u7559\u6700\u5177\u8868\u73b0\u529b\u7684\u8f93\u51fa\n        if self.glu_adapter is not None:\n            hidden_states = self.glu_adapter(hidden_states)\n        # hidden_states = self.glu_adapter(hidden_states)\n        hidden_states = attention_out + hidden_states\n        \n        return hidden_states, residual\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        return self.mixer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype, **kwargs)\n\n\ndef create_block(\n    d_model,\n    ssm_cfg=None,\n    norm_epsilon=1e-5,\n    drop_path=0.,\n    rms_norm=False,\n    input_resolution=(14, 14), \n    num_heads=8,\n    residual_in_fp32=False,\n    fused_add_norm=False,\n    layer_idx=None,\n    integrate_glu=False,  # \u65b0\u589e\u53c2\u6570\n    device=None,\n    dtype=None,\n    if_bimamba=False,\n    bimamba_type=\"none\",\n    if_devide_out=False,\n    init_layer_scale=None,\n):\n    if if_bimamba:\n        bimamba_type = \"v1\"\n    if ssm_cfg is None:\n        ssm_cfg = {}\n    factory_kwargs = {\"device\": device, \"dtype\": dtype}\n    mixer_cls = partial(Mamba, layer_idx=layer_idx, bimamba_type=bimamba_type, if_devide_out=if_devide_out, init_layer_scale=init_layer_scale, **ssm_cfg, **factory_kwargs)\n    norm_cls = partial(\n        nn.LayerNorm if not rms_norm else RMSNorm, eps=norm_epsilon, **factory_kwargs\n    )\n    block = Block(\n        d_model,\n        mixer_cls,\n        attention_cls=LinearAttention,\n        input_resolution= input_resolution,\n        num_heads= num_heads,\n        norm_cls=norm_cls,\n        drop_path=drop_path,\n        fused_add_norm=fused_add_norm,\n        residual_in_fp32=residual_in_fp32,\n        integrate_glu=integrate_glu,  # \u4f20\u9012\u53c2\u6570\n    )\n    block.layer_idx = layer_idx\n    return block\n\n\n# https://github.com/huggingface/transformers/blob/c28d04e9e252a1a099944e325685f14d242ecdcd/src/transformers/models/gpt2/modeling_gpt2.py#L454\ndef _init_weights(\n    module,\n    n_layer,\n    initializer_range=0.02,  # Now only used for embedding layer.\n    rescale_prenorm_residual=True,\n    n_residuals_per_layer=1,  # Change to 2 if we have MLP\n):\n    if isinstance(module, nn.Linear):\n        if module.bias is not None:\n            if not getattr(module.bias, \"_no_reinit\", False):\n                nn.init.zeros_(module.bias)\n    elif isinstance(module, nn.Embedding):\n        nn.init.normal_(module.weight, std=initializer_range)\n\n    if rescale_prenorm_residual:\n        for name, p in module.named_parameters():\n            if name in [\"out_proj.weight\", \"fc2.weight\"]:\n                nn.init.kaiming_uniform_(p, a=math.sqrt(5))\n                with torch.no_grad():\n                    p /= math.sqrt(n_residuals_per_layer * n_layer)\n\n\ndef segm_init_weights(m):\n    if isinstance(m, nn.Linear):\n        trunc_normal_(m.weight, std=0.02)\n        if isinstance(m, nn.Linear) and m.bias is not None:\n            nn.init.constant_(m.bias, 0)\n    elif isinstance(m, nn.Conv2d):\n        # NOTE conv was left to pytorch default in my original init\n        lecun_normal_(m.weight)\n        if m.bias is not None:\n            nn.init.zeros_(m.bias)\n    elif isinstance(m, (nn.LayerNorm, nn.GroupNorm, nn.BatchNorm2d)):\n        nn.init.zeros_(m.bias)\n        nn.init.ones_(m.weight)\n\n\nclass VisionMamba(nn.Module):\n    def __init__(self, \n                 img_size=224, \n                 patch_size=16, \n                 stride=16,\n                 depth=24, \n                 embed_dim=192, \n                 channels=3, \n                 num_classes=1000,\n                 ssm_cfg=None, \n                 drop_rate=0.,\n                 drop_path_rate=0.1,\n                 norm_epsilon: float = 1e-5, \n                 rms_norm: bool = False, \n                 initializer_cfg=None,\n                 fused_add_norm=False,\n                 residual_in_fp32=False,\n                 device=None,\n                 dtype=None,\n                 ft_seq_len=None,\n                 pt_hw_seq_len=14,\n                 if_bidirectional=False,\n                 final_pool_type='none',\n                 if_abs_pos_embed=False,\n                 if_rope=False,\n                 if_rope_residual=False,\n                 flip_img_sequences_ratio=-1.,\n                 if_bimamba=False,\n                 bimamba_type=\"none\",\n                 if_cls_token=False,\n                 if_devide_out=False,\n                 init_layer_scale=None,\n                 use_double_cls_token=False,\n                 use_middle_cls_token=False,\n                 **kwargs):\n        factory_kwargs = {\"device\": device, \"dtype\": dtype}\n        # add factory_kwargs into kwargs\n        kwargs.update(factory_kwargs) \n        super().__init__()\n        self.residual_in_fp32 = residual_in_fp32\n        self.fused_add_norm = fused_add_norm\n        self.if_bidirectional = if_bidirectional\n        self.final_pool_type = final_pool_type\n        self.if_abs_pos_embed = if_abs_pos_embed\n        self.if_rope = if_rope\n        self.if_rope_residual = if_rope_residual\n        self.flip_img_sequences_ratio = flip_img_sequences_ratio\n        # self.if_cls_token = if_cls_token\n        self.if_cls_token = False\n        self.use_double_cls_token = use_double_cls_token\n        self.use_middle_cls_token = use_middle_cls_token\n        self.num_tokens = 1 if if_cls_token else 0\n\n        # pretrain parameters\n        self.num_classes = num_classes\n        self.d_model = self.num_features = self.embed_dim = embed_dim  # num_features for consistency with other models\n        base_filter=32\n        self.hin_block = nn.Sequential(\n            nn.Conv2d(channels, base_filter, kernel_size=3, stride=1, padding=1),\n            HinResBlock(base_filter, base_filter),  # \u7b2c\u4e00\u4e2a HinResBlock\n            HinResBlock(base_filter, base_filter),  # \u7b2c\u4e8c\u4e2a HinResBlock\n            HinResBlock(base_filter, base_filter),  # \u7b2c\u4e09\u4e2a HinResBlock\n            nn.Conv2d(base_filter, channels, kernel_size=1)  # \u5c06\u901a\u9053\u6570\u4ece32\u8f6c\u6362\u4e3a3\n        )\n        ##############################\n        self.patch_embed = PatchEmbed(\n            img_size=img_size, patch_size=patch_size, stride=stride, in_chans=channels, embed_dim=embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        if if_cls_token:\n            if use_double_cls_token:\n                self.cls_token_head = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                self.cls_token_tail = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                self.num_tokens = 2\n            else:\n                self.cls_token = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n                # self.num_tokens = 1\n            \n        if if_abs_pos_embed:\n            # self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + self.num_tokens, self.embed_dim))\n            self.pos_embed_x = nn.Parameter(torch.zeros(1, 256, self.embed_dim))\n            self.pos_embed_z = nn.Parameter(torch.zeros(1, 64, self.embed_dim))\n            self.pos_drop = nn.Dropout(p=drop_rate)\n\n        if if_rope:\n            half_head_dim = embed_dim // 2\n            hw_seq_len = img_size // patch_size\n            self.rope = VisionRotaryEmbeddingFast(\n                dim=half_head_dim,\n                pt_seq_len=pt_hw_seq_len,\n                ft_seq_len=hw_seq_len\n            )\n            \n        # TODO: release this comment\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]  # stochastic depth decay rule\n        # import ipdb;ipdb.set_trace()\n        inter_dpr = [0.0] + dpr\n        self.drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()\n                # transformer blocks\n        self.layers = nn.ModuleList(\n            [\n                create_block(\n                    embed_dim,\n                    ssm_cfg=ssm_cfg,\n                    norm_epsilon=norm_epsilon,\n                    rms_norm=rms_norm,\n                    residual_in_fp32=residual_in_fp32,\n                    fused_add_norm=fused_add_norm,\n                    layer_idx=i,\n                    integrate_glu=(i == 0),  # \u4ec5\u5728\u7b2c\u4e00\u4e2a\u5757\u4e2d\u96c6\u6210 ResGLUAdapter\n                    if_bimamba=if_bimamba,\n                    bimamba_type=bimamba_type,\n                    drop_path=inter_dpr[i],\n                    if_devide_out=if_devide_out,\n                    init_layer_scale=init_layer_scale,\n                    **factory_kwargs,\n                )\n                for i in range(depth)\n            ]\n        )\n        \n        # output head\n        self.norm_f = (nn.LayerNorm if not rms_norm else RMSNorm)(\n            embed_dim, eps=norm_epsilon, **factory_kwargs\n        )\n\n        # self.pre_logits = nn.Identity()\n\n        # original init\n        self.patch_embed.apply(segm_init_weights)\n        # self.head.apply(segm_init_weights)\n        if if_abs_pos_embed:\n            trunc_normal_(self.pos_embed_x, std=.02)\n            trunc_normal_(self.pos_embed_z, std=.02)\n            \n        if if_cls_token:\n            if use_double_cls_token:\n                trunc_normal_(self.cls_token_head, std=.02)\n                trunc_normal_(self.cls_token_tail, std=.02)\n            else:\n                trunc_normal_(self.cls_token, std=.02)\n\n        # mamba init\n        self.apply(\n            partial(\n                _init_weights,\n                n_layer=depth,\n                **(initializer_cfg if initializer_cfg is not None else {}),\n            )\n        )\n\n\n    def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None, **kwargs):\n        return {\n            i: layer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype, **kwargs)\n            for i, layer in enumerate(self.layers)\n        }\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {\"pos_embed\", \"cls_token\", \"dist_token\", \"cls_token_head\", \"cls_token_tail\"}\n\n    @torch.jit.ignore()\n    def load_pretrained(self, checkpoint_path, prefix=\"\"):\n        _load_weights(self, checkpoint_path, prefix)\n\n    def forward_features(self, z, x, inference_params=None, if_random_cls_token_position=False, if_random_token_rank=False, target_layers=None):\n        # print(f\"Shape of x before hin_block: {x.shape if x is not None else 'None'}\")\n        x = self.hin_block(x)\n        z = self.hin_block(z)\n\n        x = self.patch_embed(x)                  #x.shape = torch.Size([B, 3, 256, 256])  -> torch.Size([2, 256, 384])\n        z = self.patch_embed(z)                  #z.shape = torch.Size([B, 3, 128, 128])  -> torch.Size([2, 64, 384])\n        B, M, _ = x.shape\n       \n        if self.if_cls_token:                 # False\n            if self.use_double_cls_token:\n                cls_token_head = self.cls_token_head.expand(B, -1, -1)\n                cls_token_tail = self.cls_token_tail.expand(B, -1, -1)\n                token_position = [0, M + 1]\n                x = torch.cat((cls_token_head, x, cls_token_tail), dim=1)\n                M = x.shape[1]\n            else:\n                if self.use_middle_cls_token:\n                    cls_token = self.cls_token.expand(B, -1, -1)\n                    token_position = M // 2\n                    # add cls token in the middle\n                    x = torch.cat((x[:, :token_position, :], cls_token, x[:, token_position:, :]), dim=1)       \n                elif if_random_cls_token_position:\n                    cls_token = self.cls_token.expand(B, -1, -1)\n                    token_position = random.randint(0, M)\n                    x = torch.cat((x[:, :token_position, :], cls_token, x[:, token_position:, :]), dim=1)\n                    print(\"token_position: \", token_position)\n                else:\n                    cls_token = self.cls_token.expand(B, -1, -1)  # stole cls_tokens impl from Phil Wang, thanks\n                    token_position = 0\n                    x = torch.cat((cls_token, x), dim=1)\n                M = x.shape[1]                 \n       \n        if self.if_abs_pos_embed:                  # True \n            x = x + self.pos_embed_x               # x = x + positon_embemding =torch.Size([B, 256, 384]) + torch.Size([1, 256, 384]) = torch.Size([B, 256, 384])\n            z = z + self.pos_embed_z               # z = z + positon_embemding =torch.Size([B, 64, 384]) + torch.Size([1, 64, 384]) = torch.Size([B, 64, 384])\n            x = torch.cat((z, x), dim=1)           # torch.Size([B, 320, 384])\n            x = self.pos_drop(x)                   # x.shape = torch.Size([B, 320, 384])\n            \n        if if_random_token_rank:                   #False\n            # \u751f\u6210\u968f\u673a shuffle \u7d22\u5f15\n            shuffle_indices = torch.randperm(M)\n\n            if isinstance(token_position, list):\n                print(\"original value: \", x[0, token_position[0], 0], x[0, token_position[1], 0])\n            else:\n                print(\"original value: \", x[0, token_position, 0])\n            print(\"original token_position: \", token_position)\n\n            # \u6267\u884c shuffle\n            x = x[:, shuffle_indices, :]\n\n            if isinstance(token_position, list):\n                # \u627e\u5230 cls token \u5728 shuffle \u4e4b\u540e\u7684\u65b0\u4f4d\u7f6e\n                new_token_position = [torch.where(shuffle_indices == token_position[i])[0].item() for i in range(len(token_position))]\n                token_position = new_token_position\n            else:\n                # \u627e\u5230 cls token \u5728 shuffle \u4e4b\u540e\u7684\u65b0\u4f4d\u7f6e\n                token_position = torch.where(shuffle_indices == token_position)[0].item()\n\n            if isinstance(token_position, list):\n                print(\"new value: \", x[0, token_position[0], 0], x[0, token_position[1], 0])\n            else:\n                print(\"new value: \", x[0, token_position, 0])\n            print(\"new token_position: \", token_position)\n\n\n        if_flip_img_sequences = False\n        if self.flip_img_sequences_ratio > 0 and (self.flip_img_sequences_ratio - random.random()) > 1e-5:        # False\n            x = x.flip([1])\n            if_flip_img_sequences = True\n\n        # mamba impl\n        residual = None\n        hidden_states = x\n        outputs = {}\n        i = None\n        if not self.if_bidirectional:                                 # True\n            for i, layer in enumerate(self.layers):\n                if if_flip_img_sequences and self.if_rope:            # False\n                    hidden_states = hidden_states.flip([1])\n                    if residual is not None:\n                        residual = residual.flip([1])\n\n                # rope about\n                if self.if_rope:                                       # False\n                    hidden_states = self.rope(hidden_states)\n                    if residual is not None and self.if_rope_residual:\n                        residual = self.rope(residual)\n\n                if if_flip_img_sequences and self.if_rope:             # False\n                    hidden_states = hidden_states.flip([1])\n                    if residual is not None:\n                        residual = residual.flip([1])\n\n                hidden_states, residual = layer(\n                    hidden_states, residual, inference_params=inference_params\n                )\n                # \u68c0\u67e5\u662f\u5426\u662f\u76ee\u6807\u5c42\n                if target_layers is not None and i in target_layers:\n                    outputs[i] = hidden_states\n                    # print(f\"Stored layer {i} output, hidden_states shape: {hidden_states.shape}\")\n\n        \n        else:             # False\n            # get two layers in a single for-loop\n            for i in range(len(self.layers) // 2):\n                if self.if_rope:\n                    hidden_states = self.rope(hidden_states)\n                    if residual is not None and self.if_rope_residual:\n                        residual = self.rope(residual)\n\n                hidden_states_f, residual_f = self.layers[i * 2](\n                    hidden_states, residual, inference_params=inference_params\n                )\n                hidden_states_b, residual_b = self.layers[i * 2 + 1](\n                    hidden_states.flip([1]), None if residual == None else residual.flip([1]), inference_params=inference_params\n                )\n                hidden_states = hidden_states_f + hidden_states_b.flip([1])\n                residual = residual_f + residual_b.flip([1])\n      \n        if not self.fused_add_norm:         #False\n            if residual is None:\n                residual = hidden_states\n            else:\n                residual = residual + self.drop_path(hidden_states)\n            hidden_states = self.norm_f(residual.to(dtype=self.norm_f.weight.dtype))\n        else:       #True\n            # Set prenorm=False here since we don't need the residual\n            fused_add_norm_fn = rms_norm_fn if isinstance(self.norm_f, RMSNorm) else layer_norm_fn\n            hidden_states = fused_add_norm_fn(                                         # hidden_states.shape = torch.Size([B, 320, 384])\n                self.drop_path(hidden_states),\n                self.norm_f.weight,\n                self.norm_f.bias,\n                eps=self.norm_f.eps,\n                residual=residual,\n                prenorm=False,\n                residual_in_fp32=self.residual_in_fp32,\n            )\n        # return only cls token if it exists\n        if self.if_cls_token:          #False\n            if self.use_double_cls_token:\n                return (hidden_states[:, token_position[0], :] + hidden_states[:, token_position[1], :]) / 2\n            else:\n                if self.use_middle_cls_token:\n                    return hidden_states[:, token_position, :]\n                elif if_random_cls_token_position:\n                    return hidden_states[:, token_position, :]\n                else:\n                    return hidden_states[:, token_position, :]\n\n        if self.final_pool_type == 'none':\n            final_output = hidden_states[:, -1, :]\n        elif self.final_pool_type == 'mean':\n            final_output = hidden_states.mean(dim=1)\n        elif self.final_pool_type == 'max':\n            final_output = hidden_states\n        elif self.final_pool_type == 'all':\n            final_output = hidden_states\n        else:\n            raise NotImplementedError\n        \n        outputs['final'] = hidden_states\n        # \u5982\u679c\u6ca1\u6709\u6307\u5b9a target_layers\uff0c\u53ea\u8fd4\u56de\u6700\u7ec8\u7279\u5f81\n        if not target_layers:\n            return final_output\n        # \u6709\u6307\u5b9a target_layers\uff0c\u8fd4\u56de\u5305\u542b\u6307\u5b9a\u5c42\u548c\u6700\u7ec8\u7279\u5f81\u7684\u5b57\u5178\n        return outputs\n\n    def forward(self, x, return_features=False, inference_params=None, if_random_cls_token_position=False, if_random_token_rank=False):\n        x = self.forward_features(x, inference_params, if_random_cls_token_position=if_random_cls_token_position, if_random_token_rank=if_random_token_rank)\n        if return_features:\n            return x\n        x = self.head(x)\n        if self.final_pool_type == 'max':\n            x = x.max(dim=1)[0]\n        return x\n\n\n@register_model\ndef vim_tiny_patch16_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, embed_dim=192, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n@register_model\ndef vim_tiny_patch16_stride8_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, stride=8, embed_dim=192, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n@register_model\ndef vim_small_patch16_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, embed_dim=384, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='all', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    \n    if pretrained:\n        checkpoint = torch.load(pretrained, map_location=\"cpu\")\n        missing_keys, unexpected_keys = model.load_state_dict(checkpoint[\"model\"], strict=False)\n        print('Load pretrained model from: ' + pretrained)\n    return model\n\n@register_model\ndef vim_small_patch16_stride8_224_bimambav2_final_pool_mean_abs_pos_embed_with_midclstok_div2(pretrained=False, **kwargs):\n    model = VisionMamba(\n        patch_size=16, stride=8, embed_dim=384, depth=24, rms_norm=True, residual_in_fp32=True, fused_add_norm=True, final_pool_type='mean', if_abs_pos_embed=True, if_rope=False, if_rope_residual=False, bimamba_type=\"v2\", if_cls_token=True, if_devide_out=True, use_middle_cls_token=True, **kwargs)\n    model.default_cfg = _cfg()\n    if pretrained:\n        checkpoint = torch.hub.load_state_dict_from_url(\n            url=\"to.do\",\n            map_location=\"cpu\", check_hash=True\n        )\n        model.load_state_dict(checkpoint[\"model\"])\n    return model\n\n\u4e2d\u5173\u4e8eLinearAttention\u7684\u96c6\u6210\u3002\u5206\u6790\u54ea\u91cc\u6709\u95ee\u9898\uff1f"}],"shard":{"index":42,"name":"turns-00042.parquet","rows":46529}}
