Efficient-Large-Model/Fast-dDrive
4167
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