Team Ai
Modelpublic

Efficient-Large-Model/Fast-dDrive

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
4likes167downloads
modeling.py3121 linesDownload Raw Back to root
1from dataclasses import dataclass2from typing import Any, Callable, Optional, Union3 4import torch5import torch.nn as nn6import torch.nn.functional as F7 8from transformers.activations import ACT2FN9from transformers.cache_utils import Cache, DynamicCache10from transformers.generation import GenerationMixin11from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask12from transformers.modeling_flash_attention_utils import FlashAttentionKwargs13from transformers.modeling_layers import GradientCheckpointingLayer14from transformers.modeling_outputs import BaseModelOutputWithPast, ModelOutput15from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update16from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel17from transformers.processing_utils import Unpack18from transformers.utils import auto_docstring, can_return_tuple, is_torchdynamo_compiling, logging19from transformers.utils.deprecation import deprecate_kwarg20from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm21from .configuration import Fast_dDriveConfig, Fast_dDriveTextConfig, Fast_dDriveVisionConfig22 23from torch.nn.attention.flex_attention import flex_attention, create_block_mask, or_masks24 25from functools import partial26import random27import math28 29# Context-parallel (CP) and Section-MoE-LoRA are research-only extensions that30# are not part of the canonical paper release. We provide no-op stubs so the31# original call sites behave as if running on a single-process, no-LoRA setup.32def is_cp_enabled():33    return False34 35 36def get_cp_rank():37    return 038 39 40def get_cp_size():41    return 142 43 44def get_cp_group():45    return None46 47 48def cp_allgather_kv(*args, **kwargs):49    raise RuntimeError("cp_allgather_kv called but context-parallel is disabled in the release build")50 51 52def rewrite_mask_mod_for_cp(mask_mod, *args, **kwargs):53    return mask_mod54 55 56def shard_doubled_sequence(x, *args, **kwargs):57    return x58 59 60def shard_position_ids(x, *args, **kwargs):61    return x62 63 64def _set_section_ids(*args, **kwargs):65    pass66 67 68def _build_section_ids_tensor(*args, **kwargs):69    return None70 71 72from .section_utils import (73    compute_section_block_idx_deep_static as _compute_section_block_idx_deep_static,74    build_deep_scaffold_sequences as _build_deep_scaffold_sequences,75    NULL_TOKEN_ID as _NULL_TOKEN_ID,76)77 78logger = logging.get_logger(__name__)79 80 81# flex_attention MUST be torch.compiled to use the efficient Triton sparse82# kernel; without it the call falls back to eager dense O(Q×KV) attention.83# TORCHDYNAMO_DISABLE=1 (often set in training scripts to avoid full-model84# compilation overhead) would silently turn @torch.compile into a no-op,85# so we temporarily re-enable dynamo for these two definitions only.86_dynamo_was_disabled = getattr(torch._dynamo.config, "disable", False)87if _dynamo_was_disabled:88    torch._dynamo.config.disable = False89 90@torch.compile()91def fused_flex_attention(q, k, v, mask=None):92    return flex_attention(q, k, v, block_mask=mask, enable_gqa=True)93 94# Compiled flex_attention for quadratic speculative decoding hot path.95# enable_gqa=True avoids repeat_interleave.96_compiled_flex_attention_quadratic = torch.compile(flex_attention)97 98if _dynamo_was_disabled:99    torch._dynamo.config.disable = True100 101def block_diff_mask(b, h, q_idx, kv_idx, block_size=None, n=None):102    """103    Constructs the specialized block diffusion attention mask for training104    composed of three masks:105    - **Block Diagonal Mask (M_BD)**: Self-attention within noised blocks106    - **Offset Block Causal Mask (M_OBC)**: Cross-attention for conditional context107    - **Block Causal Mask (M_BC)**: Attention to update x0108 109    Args:110        b, h: Batch and head indices (ignored for mask logic).111        q_idx, kv_idx: Query and Key indices.112        seq_len: Total sequence length.113        block_size: Defines the block structure.114 115    Returns:116        A boolean attention mask.117    """118    # Indicate whether token belongs to xt or x0119    x0_flag_q = (q_idx >= n)120    x0_flag_kv = (kv_idx >= n)121 122    # Compute block indices123    block_q = torch.where(x0_flag_q == 1,124                        (q_idx - n) // block_size,125                        q_idx // block_size)126    block_kv = torch.where(x0_flag_kv == 1,127                        (kv_idx - n) // block_size,128                        kv_idx // block_size)129 130    # **1. Block Diagonal Mask (M_BD) **131    block_diagonal = (block_q == block_kv) & (x0_flag_q == x0_flag_kv)132 133    # **2. Offset Block-Causal Mask (M_OBC) **134    offset_block_causal = (135    (block_q > block_kv)136    & (x0_flag_kv == 1)137    & (x0_flag_q == 0)138    )139 140    # **3. Block-Causal Mask (M_BC) **141    block_causal = (block_q >= block_kv) & (x0_flag_kv == 1) & (x0_flag_q == 1)142 143    # **4. Combine Masks **144    return block_diagonal | offset_block_causal | block_causal145 146 147def block_causal_mask(b, h, q_idx, kv_idx, block_size=None, n=None):148 149    # Indicate whether token belongs to xt or x0150    x0_flag_q = (q_idx >= n)151    x0_flag_kv = (kv_idx >= n)152 153    # Compute block indices154    block_q = torch.where(x0_flag_q == 1,155                        (q_idx - n) // block_size,156                        q_idx // block_size)157    block_kv = torch.where(x0_flag_kv == 1,158                        (kv_idx - n) // block_size,159                        kv_idx // block_size)160 161    # **1. Block Diagonal Mask (M_BD) **162    block_diagonal = (block_q == block_kv) & (x0_flag_q == x0_flag_kv)163 164    # **2. Offset Block-Causal Mask (M_OBC) **165    offset_block_causal = (166    (block_q > block_kv)167    & (x0_flag_kv == 1)168    & (x0_flag_q == 0)169    )170 171    # **3. Block-Causal Mask (M_BC) **172    block_causal = (q_idx >= kv_idx) & (x0_flag_kv == 1) & (x0_flag_q == 1)173 174    # **4. Combine Masks **175    return block_diagonal | offset_block_causal | block_causal176 177 178def hybrid_block_causal_mask_multiturn(b, h, q_idx, kv_idx, response_block_idx=None, turn_idx=None, n=None):179    """180    Multi-turn hybrid mask: Prompt uses causal, Response uses block causal.181    182    Args:183        response_block_idx: [seq_len] tensor, -1 for prompt, >=0 for response block index184        turn_idx: [seq_len] tensor, turn index for each position (0, 1, 2, ...)185        n: sequence length (half of total)186    187    Rules:188    - Each token can see all previous turns189    - Within current turn: prompt uses causal, response uses block causal190    - x_t response sees x_0: only tokens from current turn and before191    - x_0: standard causal mask192    193    Example for [prompt1, response1, prompt2, response2]:194    - prompt1 (turn 0): causal within turn 0 prompt195    - response1 (turn 0): sees prompt1 + block causal within response1196    - prompt2 (turn 1): sees all of turn 0 + causal within turn 1 prompt197    - response2 (turn 1): sees all of turn 0 + prompt2 + block causal within response2198    """199    x0_flag_q = (q_idx >= n)200    x0_flag_kv = (kv_idx >= n)201    202    pos_q = torch.where(x0_flag_q, q_idx - n, q_idx)203    pos_kv = torch.where(x0_flag_kv, kv_idx - n, kv_idx)204    205    block_q = response_block_idx[pos_q]206    block_kv = response_block_idx[pos_kv]207    turn_q = turn_idx[pos_q]208    turn_kv = turn_idx[pos_kv]209    210    is_prompt_q = (block_q < 0)211    is_prompt_kv = (block_kv < 0)212    213    # x_t region rules:214    # 1. Can see all previous turns: turn_q > turn_kv215    # 2. Within same turn, prompt: causal (turn same + is prompt + pos satisfies causal)216    # 3. Within same turn, response: sees all prompt in same turn + block causal for response217    # xt_same_turn_prompt_causal = ~x0_flag_q & ~x0_flag_kv & (turn_q == turn_kv) & is_prompt_q & (pos_q >= pos_kv)218    # xt_same_turn_response = ~x0_flag_q & ~x0_flag_kv & (turn_q == turn_kv) & ~is_prompt_q & (219    #     ~is_prompt_kv220    # )221    block_diagonal = ~x0_flag_q & ~x0_flag_kv & (turn_q == turn_kv) 222    223    # **2. Offset Block-Causal Mask (M_OBC) **224    offset_block_causal = (225        (turn_q > turn_kv) 226        & (x0_flag_kv == 1)227        & (x0_flag_q == 0)228    )229    # x_0 region: standard causal230    x0_causal = x0_flag_q & x0_flag_kv & (pos_q >= pos_kv)231    232    return (block_diagonal | 233            offset_block_causal | 234            x0_causal)235 236 237def eval_block_diff_mask(q_idx, kv_idx, block_size=None):238    # Compute block indices239    block_q = q_idx // block_size240    block_kv = kv_idx // block_size241 242    return torch.ones_like(block_q >= block_kv)243 244def eval_causal_mask(q_idx, kv_idx):245    return q_idx >= kv_idx246 247 248def eval_hybrid_block_causal_mask(q_idx, kv_idx, response_block_idx):249    """250    Inference-time hybrid block causal mask matching training's251    hybrid_block_causal_mask_multiturn pattern.252 253    For prompt tokens (block_idx == -1): standard causal mask.254    For response tokens: block-causal — can see all prompt tokens,255    bidirectional within same block, causal across blocks256    (block i can see blocks 0..i but not i+1..N).257 258    Args:259        q_idx:  [Q, 1] query position indices260        kv_idx: [1, K] key/value position indices261        response_block_idx: [seqlen] tensor, -1 for prompt, >=0 for block index262 263    Returns:264        [Q, K] boolean mask where True = can attend265    """266    block_q = response_block_idx[q_idx]    # [Q, 1]267    block_kv = response_block_idx[kv_idx]  # [1, K]268 269    is_prompt_q = (block_q < 0)270    is_prompt_kv = (block_kv < 0)271 272    # Prompt → prompt: standard causal273    prompt_causal = is_prompt_q & is_prompt_kv & (q_idx >= kv_idx)274    # Response → prompt: can see all prompt tokens275    response_sees_prompt = ~is_prompt_q & is_prompt_kv276    # Response → response: block causal (same block = bidirectional, earlier block = OK)277    response_block_causal = ~is_prompt_q & ~is_prompt_kv & (block_q >= block_kv)278 279    return prompt_causal | response_sees_prompt | response_block_causal280 281 282def _crop_dynamic_cache(past_key_values: DynamicCache, max_length: int):283    """Crop DynamicCache to max_length (used after draft_only phase in quadratic speculative decoding)."""284    new_past = []285    for layer_num in range(len(past_key_values)):286        layer_kv = ()287        for kv_idx in range(len(past_key_values[layer_num])):288            layer_kv += (past_key_values[layer_num][kv_idx][:, :, :max_length, :],)289        new_past.append(layer_kv)290    return DynamicCache(new_past)291 292 293def _extract_draft_kv_cache(past_key_values: DynamicCache, clean_len: int, block_length: int):294    """After quadratic decoding, extract only draft tokens (first of each block) from cache."""295    new_past = []296    for layer_num in range(len(past_key_values)):297        layer_kv = ()298        for kv_idx in range(len(past_key_values[layer_num])):299            tensor = past_key_values[layer_num][kv_idx]300            clean_part = tensor[:, :, :clean_len, :]301            draft_part = tensor[:, :, clean_len:: (block_length + 1), :]302            layer_kv += (torch.cat([clean_part, draft_part], dim=2),)303        new_past.append(layer_kv)304    return DynamicCache(new_past)305 306 307class Fast_dDriveMLP(nn.Module):308    def __init__(self, config, bias: bool = False):309        super().__init__()310        self.hidden_size = config.hidden_size311        self.intermediate_size = config.intermediate_size312        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=bias)313        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=bias)314        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=bias)315        self.act_fn = ACT2FN[config.hidden_act]316 317    def forward(self, hidden_state):318        return self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state))319 320 321class Fast_dDriveVisionPatchEmbed(nn.Module):322    def __init__(323        self,324        patch_size: int = 14,325        temporal_patch_size: int = 2,326        in_channels: int = 3,327        embed_dim: int = 1152,328    ) -> None:329        super().__init__()330        self.patch_size = patch_size331        self.temporal_patch_size = temporal_patch_size332        self.in_channels = in_channels333        self.embed_dim = embed_dim334 335        kernel_size = [temporal_patch_size, patch_size, patch_size]336        self.proj = nn.Conv3d(in_channels, embed_dim, kernel_size=kernel_size, stride=kernel_size, bias=False)337 338    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:339        target_dtype = self.proj.weight.dtype340        hidden_states = hidden_states.view(341            -1, self.in_channels, self.temporal_patch_size, self.patch_size, self.patch_size342        )343        hidden_states = self.proj(hidden_states.to(dtype=target_dtype)).view(-1, self.embed_dim)344        return hidden_states345 346 347class Fast_dDriveVisionRotaryEmbedding(nn.Module):348    inv_freq: torch.Tensor  # fix linting for `register_buffer`349 350    def __init__(self, dim: int, theta: float = 10000.0) -> None:351        super().__init__()352        inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))353        self.register_buffer("inv_freq", inv_freq, persistent=False)354 355    def forward(self, seqlen: int) -> torch.Tensor:356        seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype)357        freqs = torch.outer(seq, self.inv_freq)358        return freqs359 360 361class Fast_dDrivePatchMerger(nn.Module):362    def __init__(self, dim: int, context_dim: int, spatial_merge_size: int = 2) -> None:363        super().__init__()364        self.hidden_size = context_dim * (spatial_merge_size**2)365        self.ln_q = Qwen2RMSNorm(context_dim, eps=1e-6)366        self.mlp = nn.Sequential(367            nn.Linear(self.hidden_size, self.hidden_size),368            nn.GELU(),369            nn.Linear(self.hidden_size, dim),370        )371 372    def forward(self, x: torch.Tensor) -> torch.Tensor:373        x = self.mlp(self.ln_q(x).view(-1, self.hidden_size))374        return x375 376 377def rotate_half(x):378    """Rotates half the hidden dims of the input."""379    x1 = x[..., : x.shape[-1] // 2]380    x2 = x[..., x.shape[-1] // 2 :]381    return torch.cat((-x2, x1), dim=-1)382 383 384def apply_rotary_pos_emb_vision(385    q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor386) -> tuple[torch.Tensor, torch.Tensor]:387    orig_q_dtype = q.dtype388    orig_k_dtype = k.dtype389    q, k = q.float(), k.float()390    cos, sin = cos.unsqueeze(-2).float(), sin.unsqueeze(-2).float()391    q_embed = (q * cos) + (rotate_half(q) * sin)392    k_embed = (k * cos) + (rotate_half(k) * sin)393    q_embed = q_embed.to(orig_q_dtype)394    k_embed = k_embed.to(orig_k_dtype)395    return q_embed, k_embed396 397 398def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:399    """400    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,401    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)402    """403    batch, num_key_value_heads, slen, head_dim = hidden_states.shape404    if n_rep == 1:405        return hidden_states406    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)407    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)408 409 410def eager_attention_forward(411    module: nn.Module,412    query: torch.Tensor,413    key: torch.Tensor,414    value: torch.Tensor,415    attention_mask: Optional[torch.Tensor],416    scaling: float,417    dropout: float = 0.0,418    **kwargs,419):420    key_states = repeat_kv(key, module.num_key_value_groups)421    value_states = repeat_kv(value, module.num_key_value_groups)422 423    attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling424    if attention_mask is not None:425        causal_mask = attention_mask[:, :, :, : key_states.shape[-2]]426        attn_weights = attn_weights + causal_mask427 428    attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)429    attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)430    attn_output = torch.matmul(attn_weights, value_states)431    attn_output = attn_output.transpose(1, 2).contiguous()432 433    return attn_output, attn_weights434 435 436class Fast_dDriveVisionAttention(nn.Module):437    def __init__(self, config: Fast_dDriveVisionConfig) -> None:438        super().__init__()439        self.dim = config.hidden_size440        self.num_heads = config.num_heads441        self.head_dim = self.dim // self.num_heads442        self.num_key_value_groups = 1  # needed for eager attention443        self.qkv = nn.Linear(self.dim, self.dim * 3, bias=True)444        self.proj = nn.Linear(self.dim, self.dim)445        self.scaling = self.head_dim**-0.5446        self.config = config447        self.attention_dropout = 0.0448        self.is_causal = False449 450    def forward(451        self,452        hidden_states: torch.Tensor,453        cu_seqlens: torch.Tensor,454        rotary_pos_emb: Optional[torch.Tensor] = None,455        position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,456        **kwargs,457    ) -> torch.Tensor:458        seq_length = hidden_states.shape[0]459        query_states, key_states, value_states = (460            self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0)461        )462        cos, sin = position_embeddings463        query_states, key_states = apply_rotary_pos_emb_vision(query_states, key_states, cos, sin)464 465        query_states = query_states.transpose(0, 1).unsqueeze(0)466        key_states = key_states.transpose(0, 1).unsqueeze(0)467        value_states = value_states.transpose(0, 1).unsqueeze(0)468 469        attention_interface: Callable = eager_attention_forward470        if self.config._attn_implementation != "eager":471            attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]472 473        if self.config._attn_implementation == "flash_attention_2":474            # Flash Attention 2: Use cu_seqlens for variable length attention475            max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max()476            attn_output, _ = attention_interface(477                self,478                query_states,479                key_states,480                value_states,481                attention_mask=None,482                scaling=self.scaling,483                dropout=0.0 if not self.training else self.attention_dropout,484                cu_seq_lens_q=cu_seqlens,485                cu_seq_lens_k=cu_seqlens,486                max_length_q=max_seqlen,487                max_length_k=max_seqlen,488                is_causal=False,489                **kwargs,490            )491        else:492            # Other implementations: Process each chunk separately493            lengths = cu_seqlens[1:] - cu_seqlens[:-1]494            splits = [495                torch.split(tensor, lengths.tolist(), dim=2) for tensor in (query_states, key_states, value_states)496            ]497 498            attn_outputs = [499                attention_interface(500                    self,501                    q,502                    k,503                    v,504                    attention_mask=None,505                    scaling=self.scaling,506                    dropout=0.0 if not self.training else self.attention_dropout,507                    is_causal=False,508                    **kwargs,509                )[0]510                for q, k, v in zip(*splits)511            ]512            attn_output = torch.cat(attn_outputs, dim=1)513 514        attn_output = attn_output.reshape(seq_length, -1).contiguous()515        attn_output = self.proj(attn_output)516        return attn_output517 518 519class Fast_dDriveVisionBlock(GradientCheckpointingLayer):520    def __init__(self, config, attn_implementation: str = "sdpa") -> None:521        super().__init__()522        self.norm1 = Qwen2RMSNorm(config.hidden_size, eps=1e-6)523        self.norm2 = Qwen2RMSNorm(config.hidden_size, eps=1e-6)524        self.attn = Fast_dDriveVisionAttention(config=config)525        self.mlp = Fast_dDriveMLP(config, bias=True)526 527    def forward(528        self,529        hidden_states: torch.Tensor,530        cu_seqlens: torch.Tensor,531        rotary_pos_emb: Optional[torch.Tensor] = None,532        position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,533        **kwargs,534    ) -> torch.Tensor:535        hidden_states = hidden_states + self.attn(536            self.norm1(hidden_states),537            cu_seqlens=cu_seqlens,538            rotary_pos_emb=rotary_pos_emb,539            position_embeddings=position_embeddings,540            **kwargs,541        )542        hidden_states = hidden_states + self.mlp(self.norm2(hidden_states))543        return hidden_states544 545 546@auto_docstring547class Fast_dDrivePreTrainedModel(PreTrainedModel):548    config: Fast_dDriveConfig549    base_model_prefix = "model"550    supports_gradient_checkpointing = True551    _no_split_modules = ["Fast_dDriveDecoderLayer", "Fast_dDriveVisionBlock"]552    _skip_keys_device_placement = "past_key_values"553    _supports_flash_attn = True554    _supports_sdpa = True555 556    _can_compile_fullgraph = True557    _supports_attention_backend = True558 559    def gradient_checkpointing_enable(560        self,561        gradient_checkpointing_kwargs: Optional[dict[str, Any]] = None,562    ) -> None:563        """564        Ensure non-reentrant checkpointing when the trainers call into Transformers'565        gradient checkpointing helper. Flash attention kernels used by MDM do not566        support reentrant checkpointing, so we request the safer path by default.567        """568        if gradient_checkpointing_kwargs is None:569            gradient_checkpointing_kwargs = {}570        else:571            gradient_checkpointing_kwargs = dict(gradient_checkpointing_kwargs)572 573        gradient_checkpointing_kwargs.setdefault("use_reentrant", False)574        super().gradient_checkpointing_enable(gradient_checkpointing_kwargs=gradient_checkpointing_kwargs)575 576 577class Fast_dDriveVisionTransformerPretrainedModel(Fast_dDrivePreTrainedModel):578    config: Fast_dDriveVisionConfig579    _no_split_modules = ["Fast_dDriveVisionBlock"]580 581    def __init__(self, config, *inputs, **kwargs) -> None:582        super().__init__(config, *inputs, **kwargs)583        self.spatial_merge_size = config.spatial_merge_size584        self.patch_size = config.patch_size585        self.fullatt_block_indexes = config.fullatt_block_indexes586        self.window_size = config.window_size587        self.spatial_merge_unit = self.spatial_merge_size * self.spatial_merge_size588 589        self.patch_embed = Fast_dDriveVisionPatchEmbed(590            patch_size=config.patch_size,591            temporal_patch_size=config.temporal_patch_size,592            in_channels=config.in_channels,593            embed_dim=config.hidden_size,594        )595 596        head_dim = config.hidden_size // config.num_heads597        self.rotary_pos_emb = Fast_dDriveVisionRotaryEmbedding(head_dim // 2)598 599        self.blocks = nn.ModuleList([Fast_dDriveVisionBlock(config) for _ in range(config.depth)])600        self.merger = Fast_dDrivePatchMerger(601            dim=config.out_hidden_size,602            context_dim=config.hidden_size,603            spatial_merge_size=config.spatial_merge_size,604        )605        self.gradient_checkpointing = False606 607    def rot_pos_emb(self, grid_thw):608        pos_ids = []609        for t, h, w in grid_thw:610            hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)611            hpos_ids = hpos_ids.reshape(612                h // self.spatial_merge_size,613                self.spatial_merge_size,614                w // self.spatial_merge_size,615                self.spatial_merge_size,616            )617            hpos_ids = hpos_ids.permute(0, 2, 1, 3)618            hpos_ids = hpos_ids.flatten()619 620            wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)621            wpos_ids = wpos_ids.reshape(622                h // self.spatial_merge_size,623                self.spatial_merge_size,624                w // self.spatial_merge_size,625                self.spatial_merge_size,626            )627            wpos_ids = wpos_ids.permute(0, 2, 1, 3)628            wpos_ids = wpos_ids.flatten()629            pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1))630        pos_ids = torch.cat(pos_ids, dim=0)631        max_grid_size = grid_thw[:, 1:].max()632        rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)633        rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)634        return rotary_pos_emb635 636    def get_window_index(self, grid_thw):637        window_index: list = []638        cu_window_seqlens: list = [0]639        window_index_id = 0640        vit_merger_window_size = self.window_size // self.spatial_merge_size // self.patch_size641 642        for grid_t, grid_h, grid_w in grid_thw:643            llm_grid_h, llm_grid_w = (644                grid_h // self.spatial_merge_size,645                grid_w // self.spatial_merge_size,646            )647            index = torch.arange(grid_t * llm_grid_h * llm_grid_w).reshape(grid_t, llm_grid_h, llm_grid_w)648            pad_h = vit_merger_window_size - llm_grid_h % vit_merger_window_size649            pad_w = vit_merger_window_size - llm_grid_w % vit_merger_window_size650            num_windows_h = (llm_grid_h + pad_h) // vit_merger_window_size651            num_windows_w = (llm_grid_w + pad_w) // vit_merger_window_size652            index_padded = F.pad(index, (0, pad_w, 0, pad_h), "constant", -100)653            index_padded = index_padded.reshape(654                grid_t,655                num_windows_h,656                vit_merger_window_size,657                num_windows_w,658                vit_merger_window_size,659            )660            index_padded = index_padded.permute(0, 1, 3, 2, 4).reshape(661                grid_t,662                num_windows_h * num_windows_w,663                vit_merger_window_size,664                vit_merger_window_size,665            )666            seqlens = (index_padded != -100).sum([2, 3]).reshape(-1)667            index_padded = index_padded.reshape(-1)668            index_new = index_padded[index_padded != -100]669            window_index.append(index_new + window_index_id)670            cu_seqlens_tmp = seqlens.cumsum(0) * self.spatial_merge_unit + cu_window_seqlens[-1]671            cu_window_seqlens.extend(cu_seqlens_tmp.tolist())672            window_index_id += (grid_t * llm_grid_h * llm_grid_w).item()673        window_index = torch.cat(window_index, dim=0)674 675        return window_index, cu_window_seqlens676 677    def forward(self, hidden_states: torch.Tensor, grid_thw: torch.Tensor, **kwargs) -> torch.Tensor:678        """679        Args:680            hidden_states (`torch.Tensor` of shape `(seq_len, hidden_size)`):681                The final hidden states of the model.682            grid_thw (`torch.Tensor` of shape `(num_images_or_videos, 3)`):683                The temporal, height and width of feature shape of each image in LLM.684 685        Returns:686            `torch.Tensor`: hidden_states.687        """688        hidden_states = self.patch_embed(hidden_states)689        rotary_pos_emb = self.rot_pos_emb(grid_thw)690        window_index, cu_window_seqlens = self.get_window_index(grid_thw)691        cu_window_seqlens = torch.tensor(692            cu_window_seqlens,693            device=hidden_states.device,694            dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,695        )696        cu_window_seqlens = torch.unique_consecutive(cu_window_seqlens)697 698        seq_len, _ = hidden_states.size()699        hidden_states = hidden_states.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)700        hidden_states = hidden_states[window_index, :, :]701        hidden_states = hidden_states.reshape(seq_len, -1)702        rotary_pos_emb = rotary_pos_emb.reshape(seq_len // self.spatial_merge_unit, self.spatial_merge_unit, -1)703        rotary_pos_emb = rotary_pos_emb[window_index, :, :]704        rotary_pos_emb = rotary_pos_emb.reshape(seq_len, -1)705        emb = torch.cat((rotary_pos_emb, rotary_pos_emb), dim=-1)706        position_embeddings = (emb.cos(), emb.sin())707 708        cu_seqlens = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]).cumsum(709            dim=0,710            # Select dtype based on the following factors:711            #  - FA2 requires that cu_seqlens_q must have dtype int32712            #  - torch.onnx.export requires that cu_seqlens_q must have same dtype as grid_thw713            # See https://github.com/huggingface/transformers/pull/34852 for more information714            dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32,715        )716        cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0)717 718        for layer_num, blk in enumerate(self.blocks):719            if layer_num in self.fullatt_block_indexes:720                cu_seqlens_now = cu_seqlens721            else:722                cu_seqlens_now = cu_window_seqlens723 724            hidden_states = blk(725                hidden_states,726                cu_seqlens=cu_seqlens_now,727                position_embeddings=position_embeddings,728                **kwargs,729            )730 731        hidden_states = self.merger(hidden_states)732        reverse_indices = torch.argsort(window_index)733        hidden_states = hidden_states[reverse_indices, :]734 735        return hidden_states736 737 738@dataclass739@auto_docstring(740    custom_intro="""741    Base class for Llava outputs, with hidden states and attentions.742    """743)744class Fast_dDriveModelOutputWithPast(ModelOutput):745    r"""746    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):747        Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape748        `(batch_size, num_heads, sequence_length, embed_size_per_head)`)749 750        Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see751        `past_key_values` input) to speed up sequential decoding.752    rope_deltas (`torch.LongTensor` of shape `(batch_size, )`, *optional*):753        The rope index difference between sequence length and multimodal rope.754    """755 756    last_hidden_state: Optional[torch.FloatTensor] = None757    past_key_values: Optional[list[torch.FloatTensor]] = None758    hidden_states: Optional[tuple[torch.FloatTensor]] = None759    attentions: Optional[tuple[torch.FloatTensor]] = None760    rope_deltas: Optional[torch.LongTensor] = None761 762 763class Fast_dDriveRotaryEmbedding(nn.Module):764    inv_freq: torch.Tensor  # fix linting for `register_buffer`765 766    def __init__(self, config: Fast_dDriveTextConfig, device=None):767        super().__init__()768        # BC: "rope_type" was originally "type"769        if hasattr(config, "rope_scaling") and config.rope_scaling is not None:770            self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type"))771        else:772            self.rope_type = "default"773        self.max_seq_len_cached = config.max_position_embeddings774        self.original_max_seq_len = config.max_position_embeddings775 776        self.config = config777        self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]778 779        inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)780        self.register_buffer("inv_freq", inv_freq, persistent=False)781        self.original_inv_freq = self.inv_freq782 783    @torch.no_grad()784    @dynamic_rope_update  # power user: used with advanced RoPE types (e.g. dynamic rope)785    def forward(self, x, position_ids):786        # In contrast to other models, Fast_dDrive has different position ids for the grids787        # So we expand the inv_freq to shape (3, ...)788        inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)789        position_ids_expanded = position_ids[:, :, None, :].float()  # shape (3, bs, 1, positions)790 791        device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"792        with torch.autocast(device_type=device_type, enabled=False):  # Force float32793            freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)794            emb = torch.cat((freqs, freqs), dim=-1)795            cos = emb.cos() * self.attention_scaling796            sin = emb.sin() * self.attention_scaling797 798        return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)799 800 801class Qwen2MLP(nn.Module):802    def __init__(self, config):803        super().__init__()804        self.config = config805        self.hidden_size = config.hidden_size806        self.intermediate_size = config.intermediate_size807        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)808        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)809        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)810        self.act_fn = ACT2FN[config.hidden_act]811 812    def forward(self, x):813        down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))814        return down_proj815 816 817def apply_multimodal_rotary_pos_emb(q, k, cos, sin, mrope_section, unsqueeze_dim=1):818    """Applies Rotary Position Embedding with Multimodal Sections to the query and key tensors (https://qwenlm.github.io/blog/qwen2-vl/).819 820    Explanation:821        Multimodal 3D rotary position embedding is an extension to 1D rotary position embedding. The input embedding822        sequence contains vision (images / videos) embedding and text embedding or just contains text embedding. For823        vision embedding part, we apply rotary position embedding on temporal, height and width dimension separately.824        Here we split the channel dimension to 3 chunks for the temporal, height and width rotary position embedding.825        For text embedding part, we just apply 1D rotary position embedding. The three rotary position index (temporal,826        height and width) of text embedding is always the same, so the text embedding rotary position embedding has no827        difference with modern LLMs.828 829    Args:830        q (`torch.Tensor`): The query tensor.831        k (`torch.Tensor`): The key tensor.832        cos (`torch.Tensor`): The cosine part of the rotary embedding.833        sin (`torch.Tensor`): The sine part of the rotary embedding.834        position_ids (`torch.Tensor`):835            The position indices of the tokens corresponding to the query and key tensors. For example, this can be836            used to pass offsetted position ids when working with a KV-cache.837        mrope_section(`List(int)`):838            Multimodal rope section is for channel dimension of temporal, height and width in rope calculation.839        unsqueeze_dim (`int`, *optional*, defaults to 1):840            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and841            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note842            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and843            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes844            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have845            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.846    Returns:847        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.848    """849    mrope_section = mrope_section * 2850    cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(851        unsqueeze_dim852    )853    sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(854        unsqueeze_dim855    )856 857    q_embed = (q * cos) + (rotate_half(q) * sin)858    k_embed = (k * cos) + (rotate_half(k) * sin)859    return q_embed, k_embed860 861 862class Fast_dDriveAttention(nn.Module):863    """864    Multi-headed attention from 'Attention Is All You Need' paper. Modified to use sliding window attention: Longformer865    and "Generating Long Sequences with Sparse Transformers".866    """867 868    def __init__(self, config: Fast_dDriveTextConfig, layer_idx: Optional[int] = None):869        super().__init__()870        self.config = config871        self.layer_idx = layer_idx872        if layer_idx is None:873            logger.warning_once(874                f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will "875                "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` "876                "when creating this class."877            )878 879        self.hidden_size = config.hidden_size880        self.num_heads = config.num_attention_heads881        self.head_dim = self.hidden_size // self.num_heads882        self.num_key_value_heads = config.num_key_value_heads883        self.num_key_value_groups = self.num_heads // self.num_key_value_heads884        self.is_causal = True885        self.attention_dropout = config.attention_dropout886        self.rope_scaling = config.rope_scaling887        self.scaling = self.head_dim**-0.5888 889        if (self.head_dim * self.num_heads) != self.hidden_size:890            raise ValueError(891                f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}"892                f" and `num_heads`: {self.num_heads})."893            )894        self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=True)895        self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True)896        self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True)897        self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)898        self.sliding_window = config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None899 900        self.rotary_emb = Fast_dDriveRotaryEmbedding(config=config)901        self._quadratic_block_mask: dict = {}902 903    def _get_sbd_inference_quadratic_decoding_block_mask(self, block_length: int):904        """Build block mask for quadratic speculative decoding (one forward for full block)."""905        if block_length not in self._quadratic_block_mask:906            draft_len = block_length * (block_length + 1)907 908            def quadratic(b, h, q_idx, kv_idx):909                first_clean = torch.logical_and(910                    kv_idx % (block_length + 1) == 0,911                    kv_idx < draft_len,912                )913                first_clean = torch.logical_and(first_clean, q_idx >= kv_idx)914                block_q = q_idx // (block_length + 1)915                block_kv = kv_idx // (block_length + 1)916                same_block = torch.logical_and(block_q == block_kv, q_idx < draft_len)917                same_block_except_first = torch.logical_and(918                    same_block,919                    q_idx % (block_length + 1) != 0,920                )921                draft_part = torch.logical_or(first_clean, same_block_except_first)922                clean_part = kv_idx >= draft_len923                return torch.logical_or(draft_part, clean_part)924 925            block_mask = create_block_mask(926                quadratic,927                B=None,928                H=None,929                Q_LEN=draft_len,930                KV_LEN=draft_len + self.config.max_position_embeddings,931                device="cuda",932            )933            self._quadratic_block_mask[block_length] = block_mask934        return self._quadratic_block_mask[block_length]935 936    @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58")937    def forward(938        self,939        hidden_states: torch.Tensor,940        attention_mask: Optional[torch.Tensor] = None,941        position_ids: Optional[torch.LongTensor] = None,942        past_key_values: Optional[Cache] = None,943        output_attentions: bool = False,944        use_cache: bool = False,945        cache_position: Optional[torch.LongTensor] = None,946        position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,  # necessary, but kept here for BC947        update_kv_cache: bool = False,948        **kwargs: Unpack[FlashAttentionKwargs],949    ) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:950        bsz, q_len, _ = hidden_states.size()951 952        query_states = self.q_proj(hidden_states)953        key_states = self.k_proj(hidden_states)954        value_states = self.v_proj(hidden_states)955 956        query_states = query_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)957        key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)958        value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)959 960        cos, sin = position_embeddings961        if self.training:962            #split q into two parts963            q_1 = query_states[:,:,:query_states.shape[2]//2]964            q_2 = query_states[:,:,query_states.shape[2]//2:]965            #split k into two parts966            k_1 = key_states[:,:,:key_states.shape[2]//2]967            k_2 = key_states[:,:,key_states.shape[2]//2:]968            q_1, k_1 = apply_multimodal_rotary_pos_emb(q_1, k_1, cos, sin, self.rope_scaling["mrope_section"])969            q_2, k_2 = apply_multimodal_rotary_pos_emb(q_2, k_2, cos, sin, self.rope_scaling["mrope_section"])970            query_states = torch.cat((q_1, q_2), dim=-2)971            key_states = torch.cat((k_1, k_2), dim=-2)972        else:973            query_states, key_states = apply_multimodal_rotary_pos_emb(974                query_states, key_states, cos, sin, self.rope_scaling["mrope_section"]975            )976 977        if past_key_values is not None:978            cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}  # Specific to RoPE models979            if update_kv_cache:980                key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs)981            # elif len(past_key_values) > self.layer_idx:982            elif len(past_key_values) > self.layer_idx and past_key_values[self.layer_idx][0] is not None:983                key_states = torch.cat((past_key_values[self.layer_idx][0], key_states), dim=-2)984                value_states = torch.cat((past_key_values[self.layer_idx][1], value_states), dim=-2)985 986        self_spec_mode = getattr(self.config, "self_spec_inference_mode", None)987        block_length = getattr(self.config, "block_length", None) or getattr(self.config, "bd_size", None)988        if not self.training and self_spec_mode is not None and block_length is not None:989            if self_spec_mode == "quadratic" and past_key_values is not None:990                # HOT PATH: main loop of quadratic speculative decoding.991                # Use compiled flex_attention + enable_gqa=True (no repeat_interleave)992                # for Triton sparse kernel instead of eager O(Q×KV) dense fallback.993                seq_len = key_states.shape[2]994                draft_len = block_length * (block_length + 1)995                clean_keys = key_states[:, :, :-draft_len]996                draft_keys = key_states[:, :, -draft_len:]997                clean_values = value_states[:, :, :-draft_len]998                draft_values = value_states[:, :, -draft_len:]999                key_states = torch.cat([draft_keys, clean_keys], dim=2)1000                value_states = torch.cat([draft_values, clean_values], dim=2)1001                block_mask = self._get_sbd_inference_quadratic_decoding_block_mask(block_length)1002                block_mask.seq_lengths = (draft_len, seq_len)1003                attn_output = _compiled_flex_attention_quadratic(1004                    query_states, key_states, value_states,1005                    block_mask=block_mask, enable_gqa=True,1006                )1007            else:1008                # COLD PATH: draft_only ("default") or non-cached quadratic.1009                # Called once per sample — eager flex_attention is acceptable.1010                key_states = key_states.repeat_interleave(self.num_key_value_groups, dim=1)1011                value_states = value_states.repeat_interleave(self.num_key_value_groups, dim=1)1012                if self_spec_mode == "quadratic":1013                    # Non-cached quadratic (initial forward without past_key_values)1014                    seq_len = query_states.shape[2]1015                    draft_len = block_length * (block_length + 1)1016                    clean_len = seq_len - draft_len1017 1018                    def _causal_mask(b, h, q_idx, kv_idx):1019                        return torch.logical_and(q_idx >= kv_idx, q_idx < clean_len)1020 1021                    def _draft2clean_mask(b, h, q_idx, kv_idx):1022                        full_clean = torch.logical_and(q_idx >= clean_len, kv_idx < clean_len)1023                        first_clean = torch.logical_and(1024                            q_idx >= clean_len,1025                            (kv_idx - clean_len) % (block_length + 1) == 0,1026                        )1027                        first_clean = torch.logical_and(first_clean, q_idx >= kv_idx)1028                        return torch.logical_or(full_clean, first_clean)1029 1030                    def _draft_mask(b, h, q_idx, kv_idx):1031                        block_q = (q_idx - clean_len) // (block_length + 1)1032                        block_kv = (kv_idx - clean_len) // (block_length + 1)1033                        quadrant = torch.logical_and(q_idx >= clean_len, kv_idx >= clean_len)1034                        same_block = torch.logical_and(block_q == block_kv, quadrant)1035                        same_block_except_first = torch.logical_and(1036                            same_block,1037                            (q_idx - clean_len) % (block_length + 1) != 0,1038                        )1039                        return torch.logical_and(same_block, same_block_except_first)1040 1041                    mask = or_masks(_causal_mask, _draft2clean_mask)1042                    mask = or_masks(mask, _draft_mask)1043                    block_mask = create_block_mask(1044                        mask, B=None, H=None, Q_LEN=seq_len, KV_LEN=seq_len,1045                    )1046                else:1047                    # self_spec_mode == "default": clean causal + draft sees all1048                    seq_len = query_states.shape[2]1049                    prefix_len = seq_len - block_length1050 1051                    def _clean_q_mask(b, h, q_idx, kv_idx):1052                        return torch.logical_and(q_idx >= kv_idx, q_idx < prefix_len)1053 1054                    def _noisy_q_mask(b, h, q_idx, kv_idx):1055                        return q_idx >= prefix_len1056 1057                    block_mask = create_block_mask(1058                        or_masks(_clean_q_mask, _noisy_q_mask),1059                        B=None,1060                        H=None,1061                        Q_LEN=seq_len,1062                        KV_LEN=seq_len,1063                    )1064                attn_output = flex_attention(query_states, key_states, value_states, block_mask=block_mask)1065            attn_output = attn_output.transpose(1, 2).reshape(bsz, q_len, -1).contiguous()1066            attn_output = self.o_proj(attn_output)1067            return attn_output, None1068 1069        if self.training:1070            # CP: all-gather KV so each rank sees full sequence1071            if is_cp_enabled():1072                key_states, value_states = cp_allgather_kv(1073                    key_states.contiguous(), value_states.contiguous()1074                )1075 1076            query_states = query_states.contiguous()1077            key_states = key_states.contiguous()1078            value_states = value_states.contiguous()1079 1080            # RoPE produces fp32 Q,K (cos/sin are fp32) while V stays bf16.1081            # flex_attention requires Q,K,V to share a dtype.1082            if query_states.dtype != value_states.dtype:1083                query_states = query_states.to(value_states.dtype)1084                key_states = key_states.to(value_states.dtype)1085 1086            attn_output = fused_flex_attention(query_states, key_states, value_states, mask=attention_mask)1087            attn_output = attn_output.transpose(1, 2).contiguous()1088            attn_weights = None1089        else:1090            attention_interface: Callable = eager_attention_forward1091            if self.config._attn_implementation != "eager":1092                attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]1093 1094            attn_output, attn_weights = attention_interface(1095                self,1096                query_states,1097                key_states,1098                value_states,1099                attention_mask,1100                dropout=0.0 if not self.training else self.attention_dropout,1101                scaling=self.scaling,1102                sliding_window=self.sliding_window,1103                position_ids=position_ids,  # pass positions for FA21104                **kwargs,1105            )1106 1107        attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()1108        attn_output = self.o_proj(attn_output)1109        return attn_output, attn_weights1110 1111 1112class Fast_dDriveDecoderLayer(GradientCheckpointingLayer):1113    def __init__(self, config: Fast_dDriveTextConfig, layer_idx: int):1114        super().__init__()1115        self.hidden_size = config.hidden_size1116 1117        if config.use_sliding_window and config._attn_implementation != "flash_attention_2":1118            logger.warning_once(1119                f"Sliding Window Attention is enabled but not implemented for `{config._attn_implementation}`; "1120                "unexpected results may be encountered."1121            )1122        self.self_attn = Fast_dDriveAttention(config, layer_idx)1123 1124        self.mlp = Qwen2MLP(config)1125        self.input_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)1126        self.post_attention_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)1127        self.attention_type = config.layer_types[layer_idx]1128 1129    @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58")1130    def forward(1131        self,1132        hidden_states: torch.Tensor,1133        attention_mask: Optional[torch.Tensor] = None,1134        position_ids: Optional[torch.LongTensor] = None,1135        past_key_values: Optional[tuple[torch.Tensor]] = None,1136        output_attentions: Optional[bool] = False,1137        use_cache: Optional[bool] = False,1138        cache_position: Optional[torch.LongTensor] = None,1139        position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,  # necessary, but kept here for BC1140        update_kv_cache: bool = False,1141        **kwargs: Unpack[FlashAttentionKwargs],1142    ) -> tuple[torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]]:1143        """1144        Args:1145            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`1146            attention_mask (`torch.FloatTensor`, *optional*): attention mask of size1147                `(batch, sequence_length)` where padding elements are indicated by 0.1148            output_attentions (`bool`, *optional*):1149                Whether or not to return the attentions tensors of all attention layers. See `attentions` under1150                returned tensors for more detail.1151            use_cache (`bool`, *optional*):1152                If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding1153                (see `past_key_values`).1154            past_key_values (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states1155            cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):1156                Indices depicting the position of the input sequence tokens in the sequence.1157            position_embeddings (`tuple[torch.FloatTensor, torch.FloatTensor]`, *optional*):1158                Tuple containing the cosine and sine positional embeddings of shape `(batch_size, seq_len, head_dim)`,1159                with `head_dim` being the embedding dimension of each attention head.1160            kwargs (`dict`, *optional*):1161                Arbitrary kwargs to be ignored, used for FSDP and other methods that injects code1162                into the model1163        """1164 1165        residual = hidden_states1166 1167        hidden_states = self.input_layernorm(hidden_states)1168 1169        # Self Attention1170        hidden_states, self_attn_weights = self.self_attn(1171            hidden_states=hidden_states,1172            attention_mask=attention_mask,1173            position_ids=position_ids,1174            past_key_values=past_key_values,1175            output_attentions=output_attentions,1176            use_cache=use_cache,1177            cache_position=cache_position,1178            position_embeddings=position_embeddings,1179            update_kv_cache=update_kv_cache,1180            **kwargs,1181        )1182        hidden_states = residual + hidden_states1183 1184        # Fully Connected1185        residual = hidden_states1186        hidden_states = self.post_attention_layernorm(hidden_states)1187        hidden_states = self.mlp(hidden_states)1188        hidden_states = residual + hidden_states1189 1190        outputs = (hidden_states,)1191 1192        if output_attentions:1193            outputs += (self_attn_weights,)1194 1195        return outputs1196 1197 1198@auto_docstring1199class Fast_dDriveTextModel(Fast_dDrivePreTrainedModel):1200    config: Fast_dDriveTextConfig

Showing the first 1,200 of 3121 lines. Download the file for the rest.