| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| import functools |
| import math |
| from collections import OrderedDict |
| from abc import ABC, abstractmethod |
| import dataclasses |
| from dataclasses import dataclass |
| from typing import Any, Optional, Union |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch import Tensor |
|
|
| from ..core.attention.attention import attention_forward |
| from ..core import gradient_checkpoint_forward |
|
|
| from transformers.modeling_utils import PreTrainedModel |
| from transformers.configuration_utils import PretrainedConfig |
|
|
|
|
| class ClassInstantier(OrderedDict): |
| def __getitem__(self, key): |
| content = super().__getitem__(key) |
| cls, kwargs = content if isinstance(content, tuple) else (content, {}) |
| return cls(**kwargs) |
|
|
|
|
| class SiLUActivation(nn.Module): |
|
|
| def forward(self, input: Tensor) -> Tensor: |
| return nn.functional.silu(input) |
|
|
|
|
| class GELUTanh(nn.Module): |
| def __init__(self, use_gelu_tanh_python: bool = False): |
| super().__init__() |
| if use_gelu_tanh_python: |
| self.act = self._gelu_tanh_python |
| else: |
| self.act = functools.partial(nn.functional.gelu, approximate="tanh") |
|
|
| def _gelu_tanh_python(self, input: Tensor) -> Tensor: |
| return input * 0.5 * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (input + 0.044715 * torch.pow(input, 3.0)))) |
|
|
| def forward(self, input: Tensor) -> Tensor: |
| return self.act(input) |
|
|
|
|
| ACT2CLS = { |
| "gelu_pytorch_tanh": GELUTanh, |
| "silu": SiLUActivation, |
| } |
| ACT2FN = ClassInstantier(ACT2CLS) |
|
|
|
|
| class CacheLayerMixin(ABC): |
| """Base, abstract class for a single layer's cache.""" |
| is_compileable = False |
|
|
| def __init__(self): |
| self.keys: Optional[torch.Tensor] = None |
| self.values: Optional[torch.Tensor] = None |
| self.is_initialized = False |
|
|
| def __repr__(self): |
| return f"{self.__class__.__name__}" |
|
|
| @abstractmethod |
| def lazy_initialization(self, key_states: torch.Tensor): ... |
|
|
| @abstractmethod |
| def update(self, key_states: torch.Tensor, value_states: torch.Tensor, cache_kwargs: Optional[dict[str, Any]] = None) -> tuple[torch.Tensor, torch.Tensor]: ... |
|
|
| @abstractmethod |
| def get_mask_sizes(self, cache_position: torch.Tensor) -> tuple[int, int]: ... |
|
|
| @abstractmethod |
| def get_seq_length(self) -> int: ... |
|
|
| @abstractmethod |
| def get_max_cache_shape(self) -> int: ... |
|
|
| def offload(self): |
| if self.is_initialized: |
| self.keys = self.keys.to("cpu", non_blocking=True) |
| self.values = self.values.to("cpu", non_blocking=True) |
|
|
| def prefetch(self): |
| if self.is_initialized and self.keys.device != self.device: |
| self.keys = self.keys.to(self.device, non_blocking=True) |
| self.values = self.values.to(self.device, non_blocking=True) |
|
|
| def reset(self) -> None: |
| if self.is_initialized: |
| self.keys.zero_() |
| self.values.zero_() |
| if hasattr(self, "cumulative_length"): |
| self.cumulative_length = 0 |
|
|
| def reorder_cache(self, beam_idx: torch.LongTensor) -> None: |
| if self.get_seq_length() > 0: |
| self.keys = self.keys.index_select(0, beam_idx.to(self.keys.device)) |
| self.values = self.values.index_select(0, beam_idx.to(self.values.device)) |
|
|
|
|
| class DynamicLayer(CacheLayerMixin): |
| is_sliding = False |
|
|
| def lazy_initialization(self, key_states: torch.Tensor): |
| self.dtype, self.device = key_states.dtype, key_states.device |
| self.keys = torch.tensor([], dtype=self.dtype, device=self.device) |
| self.values = torch.tensor([], dtype=self.dtype, device=self.device) |
| self.is_initialized = True |
|
|
| def update(self, key_states: torch.Tensor, value_states: torch.Tensor, cache_kwargs: Optional[dict[str, Any]] = None) -> tuple[torch.Tensor, torch.Tensor]: |
| if not self.is_initialized: |
| self.lazy_initialization(key_states) |
| self.keys = torch.cat([self.keys, key_states], dim=-2) |
| self.values = torch.cat([self.values, value_states], dim=-2) |
| return self.keys, self.values |
|
|
| def get_mask_sizes(self, cache_position: torch.Tensor) -> tuple[int, int]: |
| kv_offset = 0 |
| query_length = cache_position.shape[0] |
| kv_length = self.get_seq_length() + query_length |
| return kv_length, kv_offset |
|
|
| def get_seq_length(self) -> int: |
| if not self.is_initialized or self.keys.numel() == 0: |
| return 0 |
| return self.keys.shape[-2] |
|
|
| def get_max_cache_shape(self) -> int: |
| return -1 |
|
|
| def crop(self, max_length: int) -> None: |
| if max_length < 0: |
| max_length = self.get_seq_length() - abs(max_length) |
| if self.get_seq_length() <= max_length: |
| return |
| self.keys = self.keys[..., :max_length, :] |
| self.values = self.values[..., :max_length, :] |
|
|
| def batch_repeat_interleave(self, repeats: int) -> None: |
| if self.get_seq_length() > 0: |
| self.keys = self.keys.repeat_interleave(repeats, dim=0) |
| self.values = self.values.repeat_interleave(repeats, dim=0) |
|
|
| def batch_select_indices(self, indices: torch.Tensor) -> None: |
| if self.get_seq_length() > 0: |
| self.keys = self.keys[indices, ...] |
| self.values = self.values[indices, ...] |
|
|
|
|
| class Cache: |
| def __init__(self, layers: Optional[list[CacheLayerMixin]] = None, layer_class_to_replicate: Optional[type[CacheLayerMixin]] = None, offloading: bool = False, offload_only_non_sliding: bool = True): |
| if layers is not None and layer_class_to_replicate is not None: |
| raise ValueError("Provide exactly one of `layers` or `layer_class_to_replicate`.") |
| if layers is None and layer_class_to_replicate is None: |
| raise ValueError("Provide exactly one of `layers` or `layer_class_to_replicate`.") |
| self.layers = layers if layers is not None else [] |
| self.layer_class_to_replicate = layer_class_to_replicate |
| self.offloading = offloading |
| if self.offloading: |
| self.only_non_sliding = offload_only_non_sliding |
| self.prefetch_stream = torch.cuda.Stream() |
|
|
| def __repr__(self): |
| return f"{self.__class__.__name__}(layers={self.layers})" |
|
|
| def update(self, key_states: torch.Tensor, value_states: torch.Tensor, layer_idx: int, cache_kwargs: Optional[dict[str, Any]] = None) -> tuple[torch.Tensor, torch.Tensor]: |
| if self.layer_class_to_replicate is not None: |
| while len(self.layers) <= layer_idx: |
| self.layers.append(self.layer_class_to_replicate()) |
| if self.offloading: |
| torch.cuda.default_stream(key_states.device).wait_stream(self.prefetch_stream) |
| self.prefetch(layer_idx + 1, self.only_non_sliding) |
| keys, values = self.layers[layer_idx].update(key_states, value_states, cache_kwargs) |
| if self.offloading: |
| self.offload(layer_idx, self.only_non_sliding) |
| return keys, values |
|
|
| def get_seq_length(self, layer_idx: int = 0) -> int: |
| if layer_idx >= len(self.layers): |
| return 0 |
| return self.layers[layer_idx].get_seq_length() |
|
|
| def get_mask_sizes(self, cache_position: torch.Tensor, layer_idx: int) -> tuple[int, int]: |
| if layer_idx >= len(self.layers): |
| return cache_position.shape[0], 0 |
| return self.layers[layer_idx].get_mask_sizes(cache_position) |
|
|
| def get_max_cache_shape(self, layer_idx: int = 0) -> int: |
| if layer_idx >= len(self.layers): |
| return -1 |
| return self.layers[layer_idx].get_max_cache_shape() |
|
|
| def reset(self): |
| for layer_idx in range(len(self.layers)): |
| self.layers[layer_idx].reset() |
|
|
| def reorder_cache(self, beam_idx: torch.LongTensor): |
| for layer_idx in range(len(self.layers)): |
| self.layers[layer_idx].reorder_cache(beam_idx) |
|
|
| def crop(self, max_length: int): |
| for layer_idx in range(len(self.layers)): |
| self.layers[layer_idx].crop(max_length) |
|
|
| def batch_repeat_interleave(self, repeats: int): |
| for layer_idx in range(len(self.layers)): |
| self.layers[layer_idx].batch_repeat_interleave(repeats) |
|
|
| def batch_select_indices(self, indices: torch.Tensor): |
| for layer_idx in range(len(self.layers)): |
| self.layers[layer_idx].batch_select_indices(indices) |
|
|
| def prefetch(self, layer_idx: int, only_non_sliding: bool = True): |
| if only_non_sliding: |
| try: |
| layer_idx = layer_idx + self.is_sliding[layer_idx:].index(False) |
| except ValueError: |
| layer_idx = self.is_sliding.index(False) |
| else: |
| layer_idx = layer_idx if layer_idx < len(self.layers) else 0 |
| with torch.cuda.stream(self.prefetch_stream): |
| self.layers[layer_idx].prefetch() |
|
|
| def offload(self, layer_idx: int, only_non_sliding: bool = True): |
| if not (only_non_sliding and self.is_sliding[layer_idx]): |
| self.layers[layer_idx].offload() |
|
|
| @property |
| def is_sliding(self) -> list[bool]: |
| return [getattr(layer, "is_sliding", False) for layer in self.layers] |
|
|
| def __getitem__(self, layer_idx: int) -> tuple[torch.Tensor, torch.Tensor]: |
| if layer_idx < len(self.layers): |
| return self.layers[layer_idx].keys, self.layers[layer_idx].values |
| else: |
| raise KeyError(f"Cache only has {len(self.layers)} layers, attempted to access layer with index {layer_idx}") |
|
|
| def __iter__(self): |
| for layer_idx in range(len(self)): |
| yield (self.layers[layer_idx].keys, self.layers[layer_idx].values) |
|
|
| def __len__(self): |
| return len(self.layers) |
|
|
| @property |
| def is_compileable(self) -> bool: |
| if len(self.layers) == 0: |
| return False |
| return all(layer.is_compileable for layer in self.layers) |
|
|
| @property |
| def is_initialized(self) -> bool: |
| return len(self.layers) > 0 and all(layer.is_initialized for layer in self.layers) |
|
|
|
|
| class ModelOutput(OrderedDict): |
| """Base class for model outputs that allows additional fields.""" |
| def __post_init__(self): |
| if dataclasses.is_dataclass(self): |
| self.__dict__.update({f.name: getattr(self, f.name) for f in dataclasses.fields(self)}) |
|
|
| def __getitem__(self, key): |
| return getattr(self, key) |
|
|
| def __setitem__(self, key, value): |
| setattr(self, key, value) |
|
|
| def __iter__(self): |
| if dataclasses.is_dataclass(self): |
| for f in dataclasses.fields(self): |
| val = getattr(self, f.name) |
| if val is not None: |
| yield val |
| else: |
| for key in self.keys(): |
| yield self[key] |
|
|
| def keys(self): |
| if dataclasses.is_dataclass(self): |
| return [f.name for f in dataclasses.fields(self) if getattr(self, f.name) is not None] |
| return list(self.__dict__.keys()) |
|
|
| def values(self): |
| if dataclasses.is_dataclass(self): |
| return [getattr(self, f.name) for f in dataclasses.fields(self) if getattr(self, f.name) is not None] |
| return list(self.__dict__.values()) |
|
|
| def items(self): |
| if dataclasses.is_dataclass(self): |
| return [(f.name, getattr(self, f.name)) for f in dataclasses.fields(self) if getattr(self, f.name) is not None] |
| return list(self.__dict__.items()) |
|
|
| def __contains__(self, key): |
| return hasattr(self, key) and getattr(self, key) is not None |
|
|
|
|
| @dataclass |
| class BaseModelOutputWithPast(ModelOutput): |
| last_hidden_state: Optional[torch.FloatTensor] = None |
| past_key_values: Optional[Cache] = None |
| hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None |
| attentions: Optional[tuple[torch.FloatTensor, ...]] = None |
|
|
|
|
| def _compute_default_rope_parameters(config, device=None, seq_len=None): |
| base = config.rope_theta |
| partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0) |
| head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads |
| dim = int(head_dim * partial_rotary_factor) |
| attention_factor = 1.0 |
| inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim)) |
| return inv_freq, attention_factor |
|
|
|
|
| ROPE_INIT_FUNCTIONS = { |
| "default": _compute_default_rope_parameters, |
| } |
|
|
|
|
| class Qwen3VLVisionConfig(PretrainedConfig): |
|
|
| model_type = "qwen3_vl" |
| base_config_key = "vision_config" |
|
|
| def __init__( |
| self, |
| depth=27, |
| hidden_size=1152, |
| hidden_act="gelu_pytorch_tanh", |
| intermediate_size=4304, |
| num_heads=16, |
| in_channels=3, |
| patch_size=16, |
| spatial_merge_size=2, |
| temporal_patch_size=2, |
| out_hidden_size=3584, |
| num_position_embeddings=2304, |
| deepstack_visual_indexes=[8, 16, 24], |
| initializer_range=0.02, |
| **kwargs, |
| ): |
| super().__init__(**kwargs) |
|
|
| self.depth = depth |
| self.hidden_size = hidden_size |
| self.hidden_act = hidden_act |
| self.intermediate_size = intermediate_size |
| self.num_heads = num_heads |
| self.in_channels = in_channels |
| self.patch_size = patch_size |
| self.spatial_merge_size = spatial_merge_size |
| self.temporal_patch_size = temporal_patch_size |
| self.out_hidden_size = out_hidden_size |
| self.num_position_embeddings = num_position_embeddings |
| self.initializer_range = initializer_range |
| self.deepstack_visual_indexes = deepstack_visual_indexes |
|
|
|
|
| class Qwen3VLTextConfig(PretrainedConfig): |
|
|
| model_type = "qwen3_vl_text" |
| base_config_key = "text_config" |
|
|
| def __init__( |
| self, |
| vocab_size=151936, |
| hidden_size=4096, |
| intermediate_size=22016, |
| num_hidden_layers=32, |
| num_attention_heads=32, |
| num_key_value_heads=32, |
| head_dim=128, |
| hidden_act="silu", |
| max_position_embeddings=128000, |
| initializer_range=0.02, |
| rms_norm_eps=1e-6, |
| use_cache=True, |
| tie_word_embeddings=False, |
| rope_theta=5000000.0, |
| rope_scaling=None, |
| attention_bias=False, |
| attention_dropout=0.0, |
| **kwargs, |
| ): |
| self.vocab_size = vocab_size |
| self.max_position_embeddings = max_position_embeddings |
| self.hidden_size = hidden_size |
| self.intermediate_size = intermediate_size |
| self.num_hidden_layers = num_hidden_layers |
| self.num_attention_heads = num_attention_heads |
| if num_key_value_heads is None: |
| num_key_value_heads = num_attention_heads |
| self.num_key_value_heads = num_key_value_heads |
| self.head_dim = head_dim |
| self.hidden_act = hidden_act |
| self.initializer_range = initializer_range |
| self.rms_norm_eps = rms_norm_eps |
| self.use_cache = use_cache |
| self.rope_theta = rope_theta |
| self.rope_scaling = rope_scaling |
| self.attention_bias = attention_bias |
| self.attention_dropout = attention_dropout |
| super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs) |
|
|
|
|
| class Qwen3VLConfig(PretrainedConfig): |
|
|
| model_type = "qwen3_vl" |
| sub_configs = {"vision_config": Qwen3VLVisionConfig, "text_config": Qwen3VLTextConfig} |
| keys_to_ignore_at_inference = ["past_key_values"] |
|
|
| def __init__( |
| self, |
| text_config=None, |
| vision_config=None, |
| image_token_id=151655, |
| video_token_id=151656, |
| vision_start_token_id=151652, |
| vision_end_token_id=151653, |
| tie_word_embeddings=False, |
| **kwargs, |
| ): |
| if isinstance(vision_config, dict): |
| self.vision_config = self.sub_configs["vision_config"](**vision_config) |
| elif vision_config is None: |
| self.vision_config = self.sub_configs["vision_config"]() |
|
|
| if isinstance(text_config, dict): |
| self.text_config = self.sub_configs["text_config"](**text_config) |
| elif text_config is None: |
| self.text_config = self.sub_configs["text_config"]() |
|
|
| self.image_token_id = image_token_id |
| self.video_token_id = video_token_id |
| self.vision_start_token_id = vision_start_token_id |
| self.vision_end_token_id = vision_end_token_id |
| super().__init__(**kwargs, tie_word_embeddings=tie_word_embeddings) |
|
|
|
|
| class Qwen3VLVisionMLP(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.hidden_size = config.hidden_size |
| self.intermediate_size = config.intermediate_size |
| self.linear_fc1 = nn.Linear(self.hidden_size, self.intermediate_size, bias=True) |
| self.linear_fc2 = nn.Linear(self.intermediate_size, self.hidden_size, bias=True) |
| self.act_fn = ACT2FN[config.hidden_act] |
|
|
| def forward(self, hidden_state): |
| return self.linear_fc2(self.act_fn(self.linear_fc1(hidden_state))) |
|
|
|
|
| class Qwen3VLVisionPatchEmbed(nn.Module): |
| def __init__(self, config) -> None: |
| super().__init__() |
| self.patch_size = config.patch_size |
| self.temporal_patch_size = config.temporal_patch_size |
| self.in_channels = config.in_channels |
| self.embed_dim = config.hidden_size |
|
|
| kernel_size = [self.temporal_patch_size, self.patch_size, self.patch_size] |
| self.proj = nn.Conv3d(self.in_channels, self.embed_dim, kernel_size=kernel_size, stride=kernel_size, bias=True) |
|
|
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| target_dtype = self.proj.weight.dtype |
| hidden_states = hidden_states.view( |
| -1, self.in_channels, self.temporal_patch_size, self.patch_size, self.patch_size |
| ) |
| hidden_states = self.proj(hidden_states.to(dtype=target_dtype)).view(-1, self.embed_dim) |
| return hidden_states |
|
|
|
|
| class Qwen3VLVisionRotaryEmbedding(nn.Module): |
| inv_freq: torch.Tensor |
|
|
| def __init__(self, dim: int, theta: float = 10000.0) -> None: |
| super().__init__() |
| inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float) / dim)) |
| self.register_buffer("inv_freq", inv_freq, persistent=False) |
|
|
| def forward(self, seqlen: int) -> torch.Tensor: |
| seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype) |
| freqs = torch.outer(seq, self.inv_freq) |
| return freqs |
|
|
| class Qwen3VLVisionPatchMerger(nn.Module): |
| def __init__(self, config, use_postshuffle_norm=False) -> None: |
| super().__init__() |
| self.hidden_size = config.hidden_size * (config.spatial_merge_size**2) |
| self.use_postshuffle_norm = use_postshuffle_norm |
| self.norm = nn.LayerNorm(self.hidden_size if use_postshuffle_norm else config.hidden_size, eps=1e-6) |
| self.linear_fc1 = nn.Linear(self.hidden_size, self.hidden_size) |
| self.act_fn = nn.GELU() |
| self.linear_fc2 = nn.Linear(self.hidden_size, config.out_hidden_size) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| x = self.norm(x.view(-1, self.hidden_size) if self.use_postshuffle_norm else x).view(-1, self.hidden_size) |
| x = self.linear_fc2(self.act_fn(self.linear_fc1(x))) |
| return x |
|
|
|
|
| def rotate_half(x): |
| """Rotates half the hidden dims of the input.""" |
| x1 = x[..., : x.shape[-1] // 2] |
| x2 = x[..., x.shape[-1] // 2 :] |
| return torch.cat((-x2, x1), dim=-1) |
|
|
|
|
| def apply_rotary_pos_emb_vision( |
| q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| orig_q_dtype = q.dtype |
| orig_k_dtype = k.dtype |
| q, k = q.float(), k.float() |
| cos, sin = cos.unsqueeze(-2).float(), sin.unsqueeze(-2).float() |
| q_embed = (q * cos) + (rotate_half(q) * sin) |
| k_embed = (k * cos) + (rotate_half(k) * sin) |
| q_embed = q_embed.to(orig_q_dtype) |
| k_embed = k_embed.to(orig_k_dtype) |
| return q_embed, k_embed |
|
|
|
|
| def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: |
| """ |
| Equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). |
| The hidden states go from (batch, num_key_value_heads, seqlen, head_dim) |
| to (batch, num_attention_heads, seqlen, head_dim) |
| """ |
| batch, num_key_value_heads, slen, head_dim = hidden_states.shape |
| if n_rep == 1: |
| return hidden_states |
| hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) |
| return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) |
|
|
|
|
| def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): |
| """Applies rotary position embeddings to query and key states.""" |
| cos = cos.unsqueeze(unsqueeze_dim) |
| sin = sin.unsqueeze(unsqueeze_dim) |
| q_embed = (q * cos) + (rotate_half(q) * sin) |
| k_embed = (k * cos) + (rotate_half(k) * sin) |
| return q_embed, k_embed |
|
|
|
|
| class Qwen3VLVisionAttention(nn.Module): |
| def __init__(self, config) -> None: |
| super().__init__() |
| self.dim = config.hidden_size |
| self.num_heads = config.num_heads |
| self.head_dim = self.dim // self.num_heads |
| self.num_key_value_groups = 1 |
| self.qkv = nn.Linear(self.dim, self.dim * 3, bias=True) |
| self.proj = nn.Linear(self.dim, self.dim) |
| self.scaling = self.head_dim**-0.5 |
| self.config = config |
| self.attention_dropout = 0.0 |
| self.is_causal = False |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| cu_seqlens: torch.Tensor, |
| rotary_pos_emb: Optional[torch.Tensor] = None, |
| position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, |
| **kwargs, |
| ) -> torch.Tensor: |
| seq_length = hidden_states.shape[0] |
| query_states, key_states, value_states = ( |
| self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0) |
| ) |
| cos, sin = position_embeddings |
| query_states, key_states = apply_rotary_pos_emb_vision(query_states, key_states, cos, sin) |
|
|
| query_states = query_states.transpose(0, 1).unsqueeze(0) |
| key_states = key_states.transpose(0, 1).unsqueeze(0) |
| value_states = value_states.transpose(0, 1).unsqueeze(0) |
|
|
| |
| lengths = cu_seqlens[1:] - cu_seqlens[:-1] |
| splits = [ |
| torch.split(tensor, lengths.tolist(), dim=2) |
| for tensor in (query_states, key_states, value_states) |
| ] |
|
|
| attn_outputs = [ |
| attention_forward( |
| q, k, v, |
| q_pattern="b n s d", k_pattern="b n s d", v_pattern="b n s d", |
| out_pattern="b n s d", |
| scale=self.scaling, |
| ) |
| for q, k, v in zip(*splits) |
| ] |
| attn_output = torch.cat(attn_outputs, dim=2) |
| attn_output = attn_output.transpose(1, 2).contiguous() |
| attn_output = attn_output.reshape(seq_length, -1).contiguous() |
| attn_output = self.proj(attn_output) |
| return attn_output |
|
|
|
|
| class Qwen3VLVisionBlock(nn.Module): |
| def __init__(self, config, attn_implementation: str = "sdpa") -> None: |
| super().__init__() |
| self.norm1 = nn.LayerNorm(config.hidden_size, eps=1e-6) |
| self.norm2 = nn.LayerNorm(config.hidden_size, eps=1e-6) |
| self.attn = Qwen3VLVisionAttention(config=config) |
| self.mlp = Qwen3VLVisionMLP(config=config) |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| cu_seqlens: torch.Tensor, |
| rotary_pos_emb: Optional[torch.Tensor] = None, |
| position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, |
| **kwargs, |
| ) -> torch.Tensor: |
| hidden_states = hidden_states + self.attn( |
| self.norm1(hidden_states), |
| cu_seqlens=cu_seqlens, |
| rotary_pos_emb=rotary_pos_emb, |
| position_embeddings=position_embeddings, |
| **kwargs, |
| ) |
| hidden_states = hidden_states + self.mlp(self.norm2(hidden_states)) |
| return hidden_states |
|
|
|
|
| class Qwen3VLTextRMSNorm(nn.Module): |
| def __init__(self, hidden_size, eps=1e-6): |
| super().__init__() |
| self.weight = nn.Parameter(torch.ones(hidden_size)) |
| self.variance_epsilon = eps |
|
|
| def forward(self, hidden_states): |
| input_dtype = hidden_states.dtype |
| hidden_states = hidden_states.float() |
| variance = hidden_states.pow(2).mean(-1, keepdim=True) |
| hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) |
| return self.weight * hidden_states.to(input_dtype) |
|
|
|
|
| def eager_attention_forward( |
| module, |
| query: torch.Tensor, |
| key: torch.Tensor, |
| value: torch.Tensor, |
| attention_mask: Optional[torch.Tensor], |
| scaling: float, |
| dropout: float = 0.0, |
| **kwargs, |
| ): |
| key_states = repeat_kv(key, module.num_key_value_groups) |
| value_states = repeat_kv(value, module.num_key_value_groups) |
|
|
| attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling |
| if attention_mask is not None: |
| causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] |
| attn_weights = attn_weights + causal_mask |
|
|
| attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) |
| attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) |
| attn_output = torch.matmul(attn_weights, value_states) |
| attn_output = attn_output.transpose(1, 2).contiguous() |
|
|
| return attn_output, attn_weights |
|
|
|
|
| class Qwen3VLTextMLP(nn.Module): |
| def __init__(self, config): |
| super().__init__() |
| self.config = config |
| self.hidden_size = config.hidden_size |
| self.intermediate_size = config.intermediate_size |
| self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) |
| self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) |
| self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) |
| self.act_fn = ACT2FN[config.hidden_act] |
|
|
| def forward(self, x): |
| down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) |
| return down_proj |
|
|
|
|
| class Qwen3VLTextDecoderLayer(nn.Module): |
| def __init__(self, config, layer_idx: int): |
| super().__init__() |
| self.hidden_size = config.hidden_size |
|
|
| self.self_attn = Qwen3VLTextAttention(config=config, layer_idx=layer_idx) |
|
|
| self.mlp = Qwen3VLTextMLP(config) |
| self.input_layernorm = Qwen3VLTextRMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| self.post_attention_layernorm = Qwen3VLTextRMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| position_embeddings: tuple[torch.Tensor, torch.Tensor], |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| past_key_values: Optional[Cache] = None, |
| use_cache: Optional[bool] = False, |
| cache_position: Optional[torch.LongTensor] = None, |
| **kwargs, |
| ) -> torch.Tensor: |
| residual = hidden_states |
| hidden_states = self.input_layernorm(hidden_states) |
| hidden_states, _ = self.self_attn( |
| hidden_states=hidden_states, |
| attention_mask=attention_mask, |
| position_ids=position_ids, |
| past_key_values=past_key_values, |
| use_cache=use_cache, |
| cache_position=cache_position, |
| position_embeddings=position_embeddings, |
| **kwargs, |
| ) |
| hidden_states = residual + hidden_states |
|
|
| residual = hidden_states |
| hidden_states = self.post_attention_layernorm(hidden_states) |
| hidden_states = self.mlp(hidden_states) |
| hidden_states = residual + hidden_states |
| return hidden_states |
|
|
|
|
| @dataclass |
| class Qwen3VLModelOutputWithPast(ModelOutput): |
| last_hidden_state: Optional[torch.FloatTensor] = None |
| past_key_values: Optional[Cache] = None |
| hidden_states: Optional[tuple[torch.FloatTensor]] = None |
| attentions: Optional[tuple[torch.FloatTensor]] = None |
| rope_deltas: Optional[torch.LongTensor] = None |
| x_pred: Optional[torch.FloatTensor] = None |
| mid_results: Optional[list] = None |
|
|
|
|
| class Qwen3VLTextAttention(nn.Module): |
| """Multi-headed attention from 'Attention Is All You Need' paper""" |
|
|
| def __init__(self, config, layer_idx: int): |
| super().__init__() |
| self.config = config |
| self.layer_idx = layer_idx |
| self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads) |
| self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads |
| self.scaling = self.head_dim**-0.5 |
| self.attention_dropout = config.attention_dropout |
| self.is_causal = True |
|
|
| self.q_proj = nn.Linear( |
| config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias |
| ) |
| self.k_proj = nn.Linear( |
| config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias |
| ) |
| self.v_proj = nn.Linear( |
| config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias |
| ) |
| self.o_proj = nn.Linear( |
| config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias |
| ) |
| self.q_norm = Qwen3VLTextRMSNorm(self.head_dim, eps=config.rms_norm_eps) |
| self.k_norm = Qwen3VLTextRMSNorm(self.head_dim, eps=config.rms_norm_eps) |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| position_embeddings: tuple[torch.Tensor, torch.Tensor], |
| attention_mask: Optional[torch.Tensor], |
| past_key_values: Optional[Cache] = None, |
| cache_position: Optional[torch.LongTensor] = None, |
| **kwargs, |
| ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: |
| input_shape = hidden_states.shape[:-1] |
| hidden_shape = (*input_shape, -1, self.head_dim) |
|
|
| query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2) |
| key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) |
| value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) |
|
|
| cos, sin = position_embeddings |
| query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) |
|
|
| if past_key_values is not None: |
| cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} |
| key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) |
|
|
| |
| key_states = repeat_kv(key_states, self.num_key_value_groups) |
| value_states = repeat_kv(value_states, self.num_key_value_groups) |
|
|
| |
| attn_output = attention_forward( |
| query_states, key_states, value_states, |
| q_pattern="b n s d", k_pattern="b n s d", v_pattern="b n s d", |
| out_pattern="b n s d", |
| attn_mask=attention_mask, |
| scale=self.scaling, |
| ) |
|
|
| attn_weights = None |
|
|
| |
| attn_output = attn_output.transpose(1, 2).flatten(2, 3).contiguous() |
| attn_output = self.o_proj(attn_output) |
| return attn_output, attn_weights |
|
|
|
|
| class BottleneckPatchEmbed(nn.Module): |
| def __init__(self, config, patch_size=32, in_chans=3, pca_dim=1024, embed_dim=4096, bias=True): |
| super().__init__() |
| self.config = config |
| self.pca_dim = pca_dim |
| self.embed_dim = embed_dim |
| self.patch_size = patch_size |
| self.in_chans = in_chans |
|
|
| self.proj1 = nn.Linear(patch_size * patch_size * in_chans, pca_dim, bias=False) |
| self.proj2 = nn.Linear(pca_dim, embed_dim, bias=bias) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| x = self.proj2(self.proj1(x)) |
| return x |
|
|
|
|
| class FinalLayer(nn.Module): |
| def __init__(self, config, hidden_size, patch_size, out_channels): |
| super().__init__() |
| self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) |
|
|
| def forward(self, x, adaln_input=None): |
| x = self.linear(x) |
| return x |
|
|
|
|
| class TimestepEmbedder(nn.Module): |
| def __init__(self, config, hidden_size, frequency_embedding_size=256): |
| super().__init__() |
| self.mlp = nn.Sequential( |
| nn.Linear(frequency_embedding_size, hidden_size, bias=True), |
| nn.SiLU(), |
| nn.Linear(hidden_size, hidden_size, bias=True), |
| ) |
| self.frequency_embedding_size = frequency_embedding_size |
|
|
| @staticmethod |
| def timestep_embedding(t, dim, max_period=10000): |
| half = dim // 2 |
| freqs = torch.exp( |
| -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half |
| ) |
| args = t.float()[:, None] * freqs[None, :] |
| embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) |
| if dim % 2: |
| embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) |
| return embedding |
|
|
| def forward(self, t): |
| t_freq = self.timestep_embedding(t * 1000, self.frequency_embedding_size) |
| t_emb = self.mlp(t_freq.to(self.mlp[0].weight.dtype)) |
| return t_emb |
|
|
|
|
| class Qwen3VLPreTrainedModel(PreTrainedModel): |
| config: Qwen3VLConfig |
| base_model_prefix = "model" |
| supports_gradient_checkpointing = True |
| _no_split_modules = ["Qwen3VLTextDecoderLayer", "Qwen3VLVisionBlock"] |
| _skip_keys_device_placement = "past_key_values" |
| _supports_flash_attn = True |
| _supports_sdpa = True |
| _can_compile_fullgraph = True |
| _supports_attention_backend = True |
| _can_record_outputs = { |
| "hidden_states": Qwen3VLTextDecoderLayer, |
| "attentions": Qwen3VLTextAttention, |
| } |
|
|
|
|
| class Qwen3VLModel(Qwen3VLPreTrainedModel): |
| config: Qwen3VLConfig |
| base_model_prefix = "model" |
| supports_gradient_checkpointing = True |
| _no_split_modules = ["Qwen3VLTextDecoderLayer", "Qwen3VLVisionBlock"] |
| _skip_keys_device_placement = "past_key_values" |
| _supports_flash_attn = True |
| _supports_sdpa = True |
| _can_compile_fullgraph = True |
| _supports_attention_backend = True |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.language_model = Qwen3VLTextModel(config.text_config) |
| self.visual = Qwen3VLVisionModel(config.vision_config) |
|
|
| self.patch_size = 32 |
| self.in_channels = 3 |
| hidden_size = config.text_config.hidden_size |
| bottleneck_dim = hidden_size // 4 |
|
|
| self.t_embedder1 = TimestepEmbedder(config, hidden_size) |
| self.x_embedder = BottleneckPatchEmbed( |
| config, |
| patch_size=self.patch_size, |
| in_chans=self.in_channels, |
| pca_dim=bottleneck_dim, |
| embed_dim=hidden_size, |
| bias=True, |
| ) |
| self.final_layer2 = FinalLayer( |
| config, |
| hidden_size=hidden_size, |
| patch_size=self.patch_size, |
| out_channels=self.in_channels, |
| ) |
|
|
| self.tms_token_id = 151673 |
| self.rope_deltas = None |
|
|
| def get_input_embeddings(self): |
| return self.language_model.get_input_embeddings() |
|
|
| def set_input_embeddings(self, value): |
| self.language_model.set_input_embeddings(value) |
|
|
| def set_decoder(self, decoder): |
| self.language_model = decoder |
|
|
| def get_decoder(self): |
| return self.language_model |
|
|
| @property |
| def language_model(self): |
| return self._language_model |
|
|
| @language_model.setter |
| def language_model(self, value): |
| self._language_model = value |
|
|
| @property |
| def visual(self): |
| return self._visual |
|
|
| @visual.setter |
| def visual(self, value): |
| self._visual = value |
|
|
| def get_video_features( |
| self, pixel_values_videos: torch.FloatTensor, video_grid_thw: Optional[torch.LongTensor] = None |
| ): |
| return self.get_image_features(pixel_values_videos, video_grid_thw) |
|
|
| def get_image_features(self, pixel_values: torch.FloatTensor, image_grid_thw: Optional[torch.LongTensor] = None): |
| pixel_values = pixel_values.type(self.visual.dtype) |
| image_embeds, deepstack_image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw) |
| split_sizes = (image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2).tolist() |
| image_embeds = torch.split(image_embeds, split_sizes) |
| return image_embeds, deepstack_image_embeds |
|
|
| def get_placeholder_mask( |
| self, |
| input_ids: torch.LongTensor, |
| inputs_embeds: torch.FloatTensor, |
| image_features: Optional[torch.FloatTensor] = None, |
| video_features: Optional[torch.FloatTensor] = None, |
| ): |
| if input_ids is None: |
| special_image_mask = inputs_embeds == self.get_input_embeddings()( |
| torch.tensor(self.config.image_token_id, dtype=torch.long, device=inputs_embeds.device) |
| ) |
| special_image_mask = special_image_mask.all(-1) |
| special_video_mask = inputs_embeds == self.get_input_embeddings()( |
| torch.tensor(self.config.video_token_id, dtype=torch.long, device=inputs_embeds.device) |
| ) |
| special_video_mask = special_video_mask.all(-1) |
| else: |
| special_image_mask = input_ids == self.config.image_token_id |
| special_video_mask = input_ids == self.config.video_token_id |
|
|
| n_image_tokens = special_image_mask.sum() |
| special_image_mask = special_image_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device) |
| if image_features is not None and inputs_embeds[special_image_mask].numel() != image_features.numel(): |
| raise ValueError( |
| f"Image features and image tokens do not match: tokens: {n_image_tokens}, features {image_features.shape[0]}" |
| ) |
|
|
| n_video_tokens = special_video_mask.sum() |
| special_video_mask = special_video_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device) |
| if video_features is not None and inputs_embeds[special_video_mask].numel() != video_features.numel(): |
| raise ValueError( |
| f"Videos features and video tokens do not match: tokens: {n_video_tokens}, features {video_features.shape[0]}" |
| ) |
|
|
| return special_image_mask, special_video_mask |
|
|
| def _run_decoder_flash(self, inputs_embeds, position_ids, token_types, return_mid_results_layers=None, |
| use_gradient_checkpointing=False, use_gradient_checkpointing_offload=False): |
| """Run decoder layers with flash attention two-pass approach. |
| |
| Replicates the Megatron attention pattern: |
| 1. Causal attention on AR tokens only (text) |
| 2. Full (bidirectional) attention on ALL tokens |
| 3. Replace AR positions with causal result (index_copy) |
| """ |
| text_model = self.language_model |
|
|
| if position_ids.ndim == 2: |
| position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) |
| elif position_ids.ndim == 3 and position_ids.shape[0] == 4: |
| position_ids = position_ids[1:] |
| position_embeddings = text_model.rotary_emb(inputs_embeds, position_ids) |
| cos, sin = position_embeddings |
|
|
| is_gen = token_types[0].bool() |
| idx_ar = torch.nonzero(~is_gen, as_tuple=False).squeeze(-1) |
|
|
| hidden_states = inputs_embeds |
| mid_results = [] if return_mid_results_layers else None |
|
|
| def _flash_layer_forward(hidden_states, decoder_layer, cos, sin, idx_ar): |
| """Flash attention layer with two-pass approach using DiffSynth attention_forward.""" |
| original_attn_forward = decoder_layer.self_attn.forward |
|
|
| def _custom_flash_attn(hidden_states, position_embeddings, attention_mask=None, **kwargs): |
| attn = decoder_layer.self_attn |
| input_shape = hidden_states.shape[:-1] |
| head_dim = attn.head_dim |
| hidden_shape = (*input_shape, -1, head_dim) |
|
|
| q = attn.q_norm(attn.q_proj(hidden_states).view(hidden_shape)) |
| k = attn.k_norm(attn.k_proj(hidden_states).view(hidden_shape)) |
| v = attn.v_proj(hidden_states).view(hidden_shape) |
|
|
| cos_pe, sin_pe = position_embeddings |
| q_r = q.transpose(1, 2) |
| k_r = k.transpose(1, 2) |
| q_r, k_r = apply_rotary_pos_emb(q_r, k_r, cos_pe, sin_pe) |
| q = q_r.transpose(1, 2).contiguous() |
| k = k_r.transpose(1, 2).contiguous() |
| v = v.contiguous() |
|
|
| softmax_scale = head_dim ** -0.5 |
|
|
| |
| q_bn = q.transpose(1, 2).contiguous() |
| k_bn = k.transpose(1, 2).contiguous() |
| v_bn = v.transpose(1, 2).contiguous() |
|
|
| |
| n_rep = attn.num_key_value_groups |
| if n_rep > 1: |
| k_bn = k_bn.repeat_interleave(n_rep, dim=1) |
| v_bn = v_bn.repeat_interleave(n_rep, dim=1) |
|
|
| |
| |
| q_ar = q_bn[:, :, idx_ar].contiguous() |
| k_ar = k_bn[:, :, idx_ar].contiguous() |
| v_ar = v_bn[:, :, idx_ar].contiguous() |
| out_ar = attention_forward( |
| q_ar, k_ar, v_ar, |
| q_pattern="b n s d", k_pattern="b n s d", v_pattern="b n s d", |
| out_pattern="b n s d", |
| is_causal=True, |
| scale=softmax_scale, |
| ) |
|
|
| |
| out_full = attention_forward( |
| q_bn, k_bn, v_bn, |
| q_pattern="b n s d", k_pattern="b n s d", v_pattern="b n s d", |
| out_pattern="b n s d", |
| is_causal=False, |
| scale=softmax_scale, |
| ) |
|
|
| |
| out_full = out_full.clone() |
| out_full[:, :, idx_ar] = out_ar |
| out_full = out_full.transpose(1, 2).contiguous() |
|
|
| attn_output = out_full.reshape(*input_shape, -1).contiguous() |
| attn_output = attn.o_proj(attn_output) |
| return attn_output, None |
|
|
| decoder_layer.self_attn.forward = _custom_flash_attn |
| try: |
| hidden_states = decoder_layer( |
| hidden_states, |
| position_embeddings=(cos, sin), |
| ) |
| finally: |
| decoder_layer.self_attn.forward = original_attn_forward |
|
|
| return hidden_states |
|
|
| for layer_idx, decoder_layer in enumerate(text_model.layers): |
| hidden_states = gradient_checkpoint_forward( |
| _flash_layer_forward, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, |
| hidden_states=hidden_states, |
| decoder_layer=decoder_layer, |
| cos=cos, |
| sin=sin, |
| idx_ar=idx_ar, |
| ) |
|
|
| if return_mid_results_layers is not None and layer_idx in return_mid_results_layers: |
| mid_results.append(hidden_states) |
|
|
| hidden_states = text_model.norm(hidden_states) |
| return hidden_states, mid_results |
|
|
| def _forward_generation(self, input_ids, position_ids, vinputs, timestep, token_types, |
| attention_mask=None, pixel_values=None, pixel_values_videos=None, |
| image_grid_thw=None, video_grid_thw=None, |
| return_mid_results_layers=None, |
| use_gradient_checkpointing=False, |
| use_gradient_checkpointing_offload=False, |
| **kwargs): |
| """Forward pass for image generation (denoising step).""" |
| inputs_embeds = self.get_input_embeddings()(input_ids) |
|
|
| if pixel_values is not None: |
| image_embeds, _ = self.get_image_features(pixel_values, image_grid_thw) |
| image_embeds = torch.cat(image_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype) |
| image_mask, _ = self.get_placeholder_mask( |
| input_ids, inputs_embeds=inputs_embeds, image_features=image_embeds) |
| inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds) |
| elif torch.is_grad_enabled(): |
| |
| |
| |
| |
| |
| |
| |
| pe = self.visual.patch_embed |
| t_sz = pe.temporal_patch_size |
| m_sz = self.visual.spatial_merge_size |
| n_patches = t_sz * m_sz * m_sz |
| patch_dim = pe.in_channels * t_sz * pe.patch_size * pe.patch_size |
| fake_pv = torch.zeros(n_patches, patch_dim, |
| device=inputs_embeds.device, |
| dtype=pe.proj.weight.dtype) |
| fake_grid = torch.tensor([[t_sz, m_sz, m_sz]], |
| dtype=torch.long, device=inputs_embeds.device) |
| fake_embs, _ = self.get_image_features(fake_pv, fake_grid) |
| fake_embs = torch.cat(fake_embs, dim=0).to(inputs_embeds.dtype) |
| inputs_embeds = inputs_embeds + fake_embs.sum() * inputs_embeds.new_zeros([]) |
|
|
| if pixel_values_videos is not None: |
| video_embeds, _ = self.get_video_features(pixel_values_videos, video_grid_thw) |
| video_embeds = torch.cat(video_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype) |
| _, video_mask = self.get_placeholder_mask( |
| input_ids, inputs_embeds=inputs_embeds, video_features=video_embeds) |
| inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds) |
|
|
| if isinstance(timestep, list): |
| timestep = torch.cat(timestep, dim=0) |
| timestep = timestep.to(inputs_embeds.device) |
| t_emb = self.t_embedder1(timestep) |
|
|
| tms_mask = (input_ids == self.tms_token_id) |
| tms_mask_3d = tms_mask.unsqueeze(-1).expand_as(inputs_embeds) |
| t_emb_expanded = t_emb.unsqueeze(1).expand_as(inputs_embeds) |
| inputs_embeds = torch.where(tms_mask_3d, t_emb_expanded, inputs_embeds) |
|
|
| if isinstance(vinputs, list): |
| vinputs = torch.cat(vinputs, dim=0) |
| vinputs = vinputs.to(inputs_embeds.device) |
| vinputs_embedded = self.x_embedder(vinputs).to(inputs_embeds.dtype) |
| inputs_embeds = torch.cat([inputs_embeds, vinputs_embedded], dim=1) |
|
|
| batch_size, total_seq_len, _ = inputs_embeds.shape |
|
|
| if isinstance(token_types, list): |
| token_types = torch.cat(token_types, dim=0) |
| token_types = token_types.to(inputs_embeds.device) |
| if token_types.dim() == 1: |
| token_types = token_types.unsqueeze(0) |
| elif token_types.dim() == 2 and token_types.shape[-1] == 1 and token_types.shape[0] == total_seq_len: |
| token_types = token_types.squeeze(-1).unsqueeze(0) |
| if token_types.shape[0] == 1 and batch_size > 1: |
| token_types = token_types.expand(batch_size, -1) |
|
|
| mid_results = None |
|
|
| hidden_states, mid_results = self._run_decoder_flash( |
| inputs_embeds, position_ids, token_types, |
| return_mid_results_layers=return_mid_results_layers, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, |
| ) |
|
|
| x_pred = self.final_layer2(hidden_states) |
|
|
| return Qwen3VLModelOutputWithPast( |
| last_hidden_state=hidden_states, |
| x_pred=x_pred, |
| mid_results=mid_results, |
| ) |
|
|
| def forward( |
| self, |
| input_ids: torch.LongTensor = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| past_key_values: Optional[Cache] = None, |
| inputs_embeds: Optional[torch.FloatTensor] = None, |
| pixel_values: Optional[torch.Tensor] = None, |
| pixel_values_videos: Optional[torch.FloatTensor] = None, |
| image_grid_thw: Optional[torch.LongTensor] = None, |
| video_grid_thw: Optional[torch.LongTensor] = None, |
| cache_position: Optional[torch.LongTensor] = None, |
| vinputs: Optional[torch.Tensor] = None, |
| timestep: Optional[torch.Tensor] = None, |
| token_types: Optional[torch.Tensor] = None, |
| return_mid_results_layers: Optional[list] = None, |
| use_gradient_checkpointing: bool = False, |
| use_gradient_checkpointing_offload: bool = False, |
| **kwargs, |
| ) -> Union[tuple, Qwen3VLModelOutputWithPast]: |
| if vinputs is not None: |
| return self._forward_generation( |
| input_ids=input_ids, position_ids=position_ids, |
| vinputs=vinputs, timestep=timestep, token_types=token_types, |
| attention_mask=attention_mask, |
| pixel_values=pixel_values, pixel_values_videos=pixel_values_videos, |
| image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, |
| return_mid_results_layers=return_mid_results_layers, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, |
| **kwargs) |
|
|
| if (input_ids is None) ^ (inputs_embeds is not None): |
| raise ValueError("You must specify exactly one of input_ids or inputs_embeds") |
|
|
| if inputs_embeds is None: |
| inputs_embeds = self.get_input_embeddings()(input_ids) |
|
|
| image_mask = None |
| video_mask = None |
|
|
| if pixel_values is not None: |
| image_embeds, deepstack_image_embeds = self.get_image_features(pixel_values, image_grid_thw) |
| image_embeds = torch.cat(image_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype) |
| image_mask, _ = self.get_placeholder_mask( |
| input_ids, inputs_embeds=inputs_embeds, image_features=image_embeds |
| ) |
| inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds) |
|
|
| if pixel_values_videos is not None: |
| video_embeds, deepstack_video_embeds = self.get_video_features(pixel_values_videos, video_grid_thw) |
| video_embeds = torch.cat(video_embeds, dim=0).to(inputs_embeds.device, inputs_embeds.dtype) |
| _, video_mask = self.get_placeholder_mask( |
| input_ids, inputs_embeds=inputs_embeds, video_features=video_embeds |
| ) |
| inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds) |
|
|
| visual_pos_masks = None |
| deepstack_visual_embeds = None |
| if image_mask is not None and video_mask is not None: |
| image_mask = image_mask[..., 0] |
| video_mask = video_mask[..., 0] |
| visual_pos_masks = image_mask | video_mask |
| deepstack_visual_embeds = [] |
| image_mask_joint = image_mask[visual_pos_masks] |
| video_mask_joint = video_mask[visual_pos_masks] |
| for img_embed, vid_embed in zip(deepstack_image_embeds, deepstack_video_embeds): |
| embed_joint = img_embed.new_zeros(visual_pos_masks.sum(), img_embed.shape[-1]).to(img_embed.device) |
| embed_joint[image_mask_joint, :] = img_embed |
| embed_joint[video_mask_joint, :] = vid_embed |
| deepstack_visual_embeds.append(embed_joint) |
| elif image_mask is not None: |
| image_mask = image_mask[..., 0] |
| visual_pos_masks = image_mask |
| deepstack_visual_embeds = deepstack_image_embeds |
| elif video_mask is not None: |
| video_mask = video_mask[..., 0] |
| visual_pos_masks = video_mask |
| deepstack_visual_embeds = deepstack_video_embeds |
|
|
| if position_ids is None: |
| attention_mask_tensor = ( |
| attention_mask if not isinstance(attention_mask, dict) else attention_mask["full_attention"] |
| ) |
| if attention_mask_tensor is not None and attention_mask_tensor.ndim == 4: |
| attention_mask_tensor = torch.diagonal(attention_mask_tensor[:, 0], dim1=1, dim2=2) |
| if attention_mask_tensor.dtype.is_floating_point: |
| attention_mask_tensor = attention_mask_tensor / torch.finfo(attention_mask_tensor.dtype).min |
| attention_mask_tensor = (1.0 - attention_mask_tensor).int() |
|
|
| if (cache_position is not None and cache_position[0] == 0) or (past_key_values is None or past_key_values.get_seq_length() == 0): |
| position_ids, rope_deltas = self.get_rope_index( |
| input_ids, image_grid_thw, video_grid_thw, attention_mask |
| ) |
| self.rope_deltas = rope_deltas |
| else: |
| q_len = inputs_embeds.shape[1] |
| position_ids = torch.arange(q_len, device=inputs_embeds.device) |
| position_ids = position_ids.view(1, -1).expand(inputs_embeds.shape[0], -1) |
| position_ids = position_ids.unsqueeze(0).expand(3, -1, -1) |
|
|
| outputs = self.language_model( |
| input_ids=None, |
| position_ids=position_ids, |
| attention_mask=attention_mask, |
| past_key_values=past_key_values, |
| inputs_embeds=inputs_embeds, |
| cache_position=cache_position, |
| visual_pos_masks=visual_pos_masks, |
| deepstack_visual_embeds=deepstack_visual_embeds, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, |
| **kwargs, |
| ) |
|
|
| return Qwen3VLModelOutputWithPast( |
| last_hidden_state=outputs.last_hidden_state, |
| past_key_values=outputs.past_key_values, |
| rope_deltas=self.rope_deltas, |
| ) |
|
|
| def get_rope_index( |
| self, |
| input_ids: Optional[torch.LongTensor] = None, |
| image_grid_thw: Optional[torch.LongTensor] = None, |
| video_grid_thw: Optional[torch.LongTensor] = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| if video_grid_thw is not None: |
| video_grid_thw = torch.repeat_interleave(video_grid_thw, video_grid_thw[:, 0], dim=0) |
| video_grid_thw[:, 0] = 1 |
|
|
| spatial_merge_size = self.config.vision_config.spatial_merge_size |
| image_token_id = self.config.image_token_id |
| video_token_id = self.config.video_token_id |
| vision_start_token_id = self.config.vision_start_token_id |
| mrope_position_deltas = [] |
| if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None): |
| total_input_ids = input_ids |
| if attention_mask is None: |
| attention_mask = torch.ones_like(total_input_ids) |
| position_ids = torch.ones( |
| 3, input_ids.shape[0], input_ids.shape[1], |
| dtype=input_ids.dtype, device=input_ids.device, |
| ) |
| image_index, video_index = 0, 0 |
| attention_mask = attention_mask.to(total_input_ids.device) |
| for i, input_ids_i in enumerate(total_input_ids): |
| input_ids_i = input_ids_i[attention_mask[i] == 1] |
| image_nums, video_nums = 0, 0 |
| vision_start_indices = torch.argwhere(input_ids_i == vision_start_token_id).squeeze(1) |
| vision_tokens = input_ids_i[vision_start_indices + 1] |
| image_nums = (vision_tokens == image_token_id).sum() |
| video_nums = (vision_tokens == video_token_id).sum() |
| input_tokens = input_ids_i.tolist() |
| llm_pos_ids_list: list = [] |
| st = 0 |
| remain_images, remain_videos = image_nums, video_nums |
| for _ in range(image_nums + video_nums): |
| if image_token_id in input_tokens and remain_images > 0: |
| ed_image = input_tokens.index(image_token_id, st) |
| else: |
| ed_image = len(input_tokens) + 1 |
| if video_token_id in input_tokens and remain_videos > 0: |
| ed_video = input_tokens.index(video_token_id, st) |
| else: |
| ed_video = len(input_tokens) + 1 |
| if ed_image < ed_video: |
| t, h, w = ( |
| image_grid_thw[image_index][0], |
| image_grid_thw[image_index][1], |
| image_grid_thw[image_index][2], |
| ) |
| image_index += 1 |
| remain_images -= 1 |
| ed = ed_image |
| else: |
| t, h, w = ( |
| video_grid_thw[video_index][0], |
| video_grid_thw[video_index][1], |
| video_grid_thw[video_index][2], |
| ) |
| video_index += 1 |
| remain_videos -= 1 |
| ed = ed_video |
| llm_grid_t, llm_grid_h, llm_grid_w = ( |
| t.item(), |
| h.item() // spatial_merge_size, |
| w.item() // spatial_merge_size, |
| ) |
| text_len = ed - st |
| st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0 |
| llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx) |
| t_index = torch.arange(llm_grid_t).view(-1, 1).expand(-1, llm_grid_h * llm_grid_w).flatten() |
| h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten() |
| w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten() |
| llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx) |
| st = ed + llm_grid_t * llm_grid_h * llm_grid_w |
|
|
| if st < len(input_tokens): |
| st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0 |
| text_len = len(input_tokens) - st |
| llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx) |
|
|
| llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1) |
| position_ids[..., i, attention_mask[i] == 1] = llm_positions.to(position_ids.device) |
| mrope_position_deltas.append(llm_positions.max() + 1 - len(total_input_ids[i])) |
| mrope_position_deltas = torch.tensor(mrope_position_deltas, device=input_ids.device).unsqueeze(1) |
| return position_ids, mrope_position_deltas |
| else: |
| if attention_mask is not None: |
| position_ids = attention_mask.long().cumsum(-1) - 1 |
| position_ids.masked_fill_(attention_mask == 0, 1) |
| position_ids = position_ids.unsqueeze(0).expand(3, -1, -1).to(attention_mask.device) |
| max_position_ids = position_ids.max(0, keepdim=False)[0].max(-1, keepdim=True)[0] |
| mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1] |
| else: |
| position_ids = ( |
| torch.arange(input_ids.shape[1], device=input_ids.device) |
| .view(1, 1, -1) |
| .expand(3, input_ids.shape[0], -1) |
| ) |
| mrope_position_deltas = torch.zeros( |
| [input_ids.shape[0], 1], device=input_ids.device, dtype=input_ids.dtype, |
| ) |
| return position_ids, mrope_position_deltas |
|
|
|
|
| class Qwen3VLTextModel(Qwen3VLPreTrainedModel): |
| config: Qwen3VLTextConfig |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.vocab_size = config.vocab_size |
|
|
| self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) |
| self.layers = nn.ModuleList( |
| [Qwen3VLTextDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] |
| ) |
| self.norm = Qwen3VLTextRMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| self.rotary_emb = Qwen3VLRotaryEmbedding(config=config) |
|
|
| self.post_init() |
|
|
| def get_input_embeddings(self): |
| return self.embed_tokens |
|
|
| def set_input_embeddings(self, value): |
| self.embed_tokens = value |
|
|
| def forward( |
| self, |
| input_ids: torch.LongTensor = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| past_key_values: Optional[Cache] = None, |
| inputs_embeds: Optional[torch.FloatTensor] = None, |
| cache_position: Optional[torch.LongTensor] = None, |
| visual_pos_masks: Optional[torch.Tensor] = None, |
| deepstack_visual_embeds: Optional[torch.FloatTensor] = None, |
| use_cache: Optional[bool] = None, |
| return_mid_results_layers: Optional[list] = None, |
| use_gradient_checkpointing: bool = False, |
| use_gradient_checkpointing_offload: bool = False, |
| **kwargs, |
| ) -> BaseModelOutputWithPast: |
| if (input_ids is None) ^ (inputs_embeds is not None): |
| raise ValueError("You must specify exactly one of input_ids or inputs_embeds") |
|
|
| if inputs_embeds is None: |
| inputs_embeds = self.embed_tokens(input_ids) |
|
|
| if use_cache is None: |
| use_cache = self.config.use_cache |
|
|
| if position_ids is None: |
| position_ids = torch.arange( |
| inputs_embeds.shape[1], device=inputs_embeds.device |
| ).unsqueeze(0).expand(inputs_embeds.shape[0], -1) |
|
|
| position_embeddings = self.rotary_emb(inputs_embeds, position_ids) |
| hidden_states = inputs_embeds |
| mid_results = [] if return_mid_results_layers else None |
|
|
| for layer_idx, decoder_layer in enumerate(self.layers): |
| hidden_states = gradient_checkpoint_forward( |
| decoder_layer, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=False, |
| hidden_states=hidden_states, |
| position_embeddings=position_embeddings, |
| attention_mask=attention_mask, |
| past_key_values=past_key_values, |
| use_cache=use_cache, |
| cache_position=cache_position, |
| ) |
|
|
| if return_mid_results_layers is not None and layer_idx in return_mid_results_layers: |
| mid_results.append(hidden_states) |
|
|
| hidden_states = self.norm(hidden_states) |
|
|
| output = BaseModelOutputWithPast( |
| last_hidden_state=hidden_states, |
| past_key_values=past_key_values, |
| ) |
| if return_mid_results_layers is not None: |
| output.mid_results = mid_results |
| return output |
|
|
|
|
| class Qwen3VLVisionModel(Qwen3VLPreTrainedModel): |
| config: Qwen3VLVisionConfig |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.spatial_merge_size = config.spatial_merge_size |
| self.patch_size = config.patch_size |
| self.spatial_merge_unit = self.spatial_merge_size * self.spatial_merge_size |
|
|
| self.patch_embed = Qwen3VLVisionPatchEmbed(config) |
|
|
| self.pos_embed = nn.Embedding(config.num_position_embeddings, config.hidden_size) |
| self.num_grid_per_side = int(config.num_position_embeddings**0.5) |
|
|
| head_dim = config.hidden_size // config.num_heads |
| self.rotary_pos_emb = Qwen3VLVisionRotaryEmbedding(head_dim // 2) |
|
|
| self.blocks = nn.ModuleList( |
| [Qwen3VLVisionBlock(config) for _ in range(config.depth)] |
| ) |
|
|
| self.merger = Qwen3VLVisionPatchMerger(config, use_postshuffle_norm=False) |
| self.deepstack_visual_indexes = config.deepstack_visual_indexes |
| self.deepstack_merger_list = nn.ModuleList( |
| [ |
| Qwen3VLVisionPatchMerger( |
| config=config, |
| use_postshuffle_norm=True, |
| ) |
| for _ in range(len(config.deepstack_visual_indexes)) |
| ] |
| ) |
| self.post_init() |
|
|
| def get_dtype(self) -> torch.dtype: |
| return self.blocks[0].mlp.linear_fc1.weight.dtype |
|
|
| def fast_pos_embed_interpolate(self, grid_thw): |
| grid_ts, grid_hs, grid_ws = grid_thw[:, 0], grid_thw[:, 1], grid_thw[:, 2] |
|
|
| idx_list = [[] for _ in range(4)] |
| weight_list = [[] for _ in range(4)] |
|
|
| for t, h, w in zip(grid_ts, grid_hs, grid_ws): |
| h_idxs = torch.linspace(0, self.num_grid_per_side - 1, h, device=self.pos_embed.weight.device) |
| w_idxs = torch.linspace(0, self.num_grid_per_side - 1, w, device=self.pos_embed.weight.device) |
|
|
| h_idxs_floor = h_idxs.int() |
| w_idxs_floor = w_idxs.int() |
| h_idxs_ceil = (h_idxs.int() + 1).clip(max=self.num_grid_per_side - 1) |
| w_idxs_ceil = (w_idxs.int() + 1).clip(max=self.num_grid_per_side - 1) |
|
|
| dh = h_idxs - h_idxs_floor |
| dw = w_idxs - w_idxs_floor |
|
|
| base_h = h_idxs_floor * self.num_grid_per_side |
| base_h_ceil = h_idxs_ceil * self.num_grid_per_side |
|
|
| indices = [ |
| (base_h[None].T + w_idxs_floor[None]).flatten(), |
| (base_h[None].T + w_idxs_ceil[None]).flatten(), |
| (base_h_ceil[None].T + w_idxs_floor[None]).flatten(), |
| (base_h_ceil[None].T + w_idxs_ceil[None]).flatten(), |
| ] |
|
|
| weights = [ |
| ((1 - dh)[None].T * (1 - dw)[None]).flatten(), |
| ((1 - dh)[None].T * dw[None]).flatten(), |
| (dh[None].T * (1 - dw)[None]).flatten(), |
| (dh[None].T * dw[None]).flatten(), |
| ] |
|
|
| for i in range(4): |
| idx_list[i].extend(indices[i].tolist()) |
| weight_list[i].extend(weights[i].tolist()) |
|
|
| idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=self.pos_embed.weight.device) |
| weight_tensor = torch.tensor( |
| weight_list, dtype=self.pos_embed.weight.dtype, device=self.pos_embed.weight.device |
| ) |
| pos_embeds = self.pos_embed(idx_tensor) * weight_tensor[:, :, None] |
| patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] |
|
|
| patch_pos_embeds = patch_pos_embeds.split([h * w for h, w in zip(grid_hs, grid_ws)]) |
|
|
| patch_pos_embeds_permute = [] |
| merge_size = self.spatial_merge_size |
| for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): |
| pos_embed = pos_embed.repeat(t, 1) |
| pos_embed = ( |
| pos_embed.view(t, h // merge_size, merge_size, w // merge_size, merge_size, -1) |
| .permute(0, 1, 3, 2, 4, 5) |
| .flatten(0, 4) |
| ) |
| patch_pos_embeds_permute.append(pos_embed) |
| patch_pos_embeds = torch.cat(patch_pos_embeds_permute) |
| return patch_pos_embeds |
|
|
| def forward( |
| self, |
| pixel_values: torch.FloatTensor, |
| grid_thw: Optional[torch.LongTensor] = None, |
| ) -> tuple[torch.Tensor, list]: |
| hidden_states = self.patch_embed(pixel_values) |
|
|
| pos_embeds = self.fast_pos_embed_interpolate(grid_thw) |
| hidden_states = hidden_states + pos_embeds |
|
|
| rotary_pos_emb = self.rot_pos_emb(grid_thw) |
|
|
| seq_len, _ = hidden_states.size() |
| hidden_states = hidden_states.reshape(seq_len, -1) |
| rotary_pos_emb = rotary_pos_emb.reshape(seq_len, -1) |
| emb = torch.cat((rotary_pos_emb, rotary_pos_emb), dim=-1) |
| position_embeddings = (emb.cos(), emb.sin()) |
|
|
| cu_seqlens = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]).cumsum( |
| dim=0, |
| dtype=grid_thw.dtype if torch.jit.is_tracing() else torch.int32, |
| ) |
| cu_seqlens = F.pad(cu_seqlens, (1, 0), value=0) |
|
|
| deepstack_feature_lists = [] |
| for layer_num, blk in enumerate(self.blocks): |
| hidden_states = blk( |
| hidden_states, |
| cu_seqlens=cu_seqlens, |
| position_embeddings=position_embeddings, |
| ) |
| if layer_num in self.deepstack_visual_indexes: |
| deepstack_feature = self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)]( |
| hidden_states |
| ) |
| deepstack_feature_lists.append(deepstack_feature) |
|
|
| hidden_states = self.merger(hidden_states) |
|
|
| return hidden_states, deepstack_feature_lists |
|
|
| def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor: |
| merge_size = self.spatial_merge_size |
|
|
| max_hw = int(grid_thw[:, 1:].max().item()) |
| freq_table = self.rotary_pos_emb(max_hw) |
| device = freq_table.device |
|
|
| total_tokens = int(torch.prod(grid_thw, dim=1).sum().item()) |
| pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device) |
|
|
| offset = 0 |
| for num_frames, height, width in grid_thw: |
| merged_h, merged_w = height // merge_size, width // merge_size |
|
|
| block_rows = torch.arange(merged_h, device=device) |
| block_cols = torch.arange(merged_w, device=device) |
| intra_row = torch.arange(merge_size, device=device) |
| intra_col = torch.arange(merge_size, device=device) |
|
|
| |
| row_idx = block_rows[:, None, None, None] * merge_size + intra_row[None, None, :, None] |
| col_idx = block_cols[None, :, None, None] * merge_size + intra_col[None, None, None, :] |
|
|
| row_idx = row_idx.expand(merged_h, merged_w, merge_size, merge_size).reshape(-1) |
| col_idx = col_idx.expand(merged_h, merged_w, merge_size, merge_size).reshape(-1) |
|
|
| coords = torch.stack((row_idx, col_idx), dim=-1) |
|
|
| if num_frames > 1: |
| coords = coords.repeat(num_frames, 1) |
|
|
| num_tokens = coords.shape[0] |
| pos_ids[offset : offset + num_tokens] = coords |
| offset += num_tokens |
|
|
| embeddings = freq_table[pos_ids] |
| embeddings = embeddings.flatten(1) |
| return embeddings |
|
|
|
|
| class Qwen3VLRotaryEmbedding(nn.Module): |
| def __init__(self, config, device=None): |
| super().__init__() |
| if hasattr(config, "rope_scaling") and config.rope_scaling is not None: |
| self.rope_type = config.rope_scaling.get("rope_type", "default") |
| else: |
| self.rope_type = "default" |
| self.max_seq_len_cached = config.max_position_embeddings |
| self.original_max_seq_len = config.max_position_embeddings |
| self.config = config |
| self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] |
| inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device) |
| self.register_buffer("inv_freq", inv_freq.detach().clone(), persistent=False) |
| self.original_inv_freq = self.inv_freq |
| self.mrope_section = config.rope_scaling.get("mrope_section", [24, 20, 20]) |
|
|
| @staticmethod |
| def compute_default_rope_parameters(config, device=None, seq_len=None): |
| return _compute_default_rope_parameters(config, device=None, seq_len=None) |
|
|
| @staticmethod |
| def apply_interleaved_mrope(freqs, mrope_section): |
| freqs_t = freqs[0] |
| for dim, offset in enumerate((1, 2), start=1): |
| length = mrope_section[dim] * 3 |
| idx = slice(offset, length, 3) |
| freqs_t[..., idx] = freqs[dim, ..., idx] |
| return freqs_t |
|
|
| def forward(self, x, position_ids): |
| if position_ids.ndim == 2: |
| position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) |
| inv_freq_expanded = self.original_inv_freq[None, None, :, None].float().to(device=x.device).expand(3, position_ids.shape[1], -1, 1) |
| position_ids_expanded = position_ids[:, :, None, :].float() |
| device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu" |
| with torch.autocast(device_type=device_type, enabled=False): |
| freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3) |
| freqs = self.apply_interleaved_mrope(freqs, self.mrope_section) |
| emb = torch.cat((freqs, freqs), dim=-1) |
| cos = emb.cos() * self.attention_scaling |
| sin = emb.sin() * self.attention_scaling |
| return cos.to(x.dtype), sin.to(x.dtype) |
|
|
|
|
| @dataclass |
| class Qwen3VLCausalLMOutputWithPast(ModelOutput): |
| loss: Optional[torch.FloatTensor] = None |
| logits: Optional[torch.FloatTensor] = None |
| past_key_values: Optional[Cache] = None |
| hidden_states: Optional[tuple[torch.FloatTensor]] = None |
| attentions: Optional[tuple[torch.FloatTensor]] = None |
| rope_deltas: Optional[torch.LongTensor] = None |
| x_pred: Optional[torch.FloatTensor] = None |
| mid_results: Optional[list] = None |
|
|
|
|
| def _build_hidream_config(): |
| text_config = Qwen3VLTextConfig( |
| hidden_size=4096, |
| num_hidden_layers=36, |
| num_attention_heads=32, |
| num_key_value_heads=8, |
| intermediate_size=12288, |
| vocab_size=151936, |
| max_position_embeddings=262144, |
| head_dim=128, |
| attention_bias=False, |
| hidden_act="silu", |
| rms_norm_eps=1e-6, |
| use_cache=True, |
| bos_token_id=151643, |
| eos_token_id=151645, |
| rope_theta=5000000, |
| attention_dropout=0.0, |
| initializer_range=0.02, |
| rope_scaling={"rope_type": "default", "mrope_section": [24, 20, 20], "mrope_interleaved": True}, |
| ) |
| vision_config = Qwen3VLVisionConfig( |
| hidden_size=1152, |
| depth=27, |
| num_heads=16, |
| intermediate_size=4304, |
| patch_size=16, |
| spatial_merge_size=2, |
| in_channels=3, |
| out_hidden_size=4096, |
| deepstack_visual_indexes=[8, 16, 24], |
| temporal_patch_size=2, |
| num_position_embeddings=2304, |
| hidden_act="gelu_pytorch_tanh", |
| initializer_range=0.02, |
| ) |
| config = Qwen3VLConfig() |
| config.text_config = text_config |
| config.vision_config = vision_config |
| config.image_token_id = 151655 |
| config.video_token_id = 151656 |
| config.vision_start_token_id = 151652 |
| config.vision_end_token_id = 151653 |
| return config |
|
|
|
|
| class HiDreamO1ImageModel(Qwen3VLPreTrainedModel): |
| _checkpoint_conversion_mapping = {} |
| _tied_weights_keys = ["lm_head.weight"] |
| accepts_loss_kwargs = False |
| config: Qwen3VLConfig |
|
|
| def __init__(self): |
| config = _build_hidream_config() |
| super().__init__(config) |
| self.model = Qwen3VLModel(config) |
| self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False) |
| self.post_init() |
|
|
| def get_input_embeddings(self): |
| return self.model.get_input_embeddings() |
|
|
| def set_input_embeddings(self, value): |
| self.model.set_input_embeddings(value) |
|
|
| def set_decoder(self, decoder): |
| self.model.set_decoder(decoder) |
|
|
| def get_decoder(self): |
| return self.model.get_decoder() |
|
|
| @property |
| def language_model(self): |
| return self.model.language_model |
|
|
| @property |
| def visual(self): |
| return self.model.visual |
|
|
| def forward( |
| self, |
| input_ids: torch.LongTensor = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| past_key_values: Optional[Cache] = None, |
| inputs_embeds: Optional[torch.FloatTensor] = None, |
| labels: Optional[torch.LongTensor] = None, |
| pixel_values: Optional[torch.Tensor] = None, |
| pixel_values_videos: Optional[torch.FloatTensor] = None, |
| image_grid_thw: Optional[torch.LongTensor] = None, |
| video_grid_thw: Optional[torch.LongTensor] = None, |
| cache_position: Optional[torch.LongTensor] = None, |
| logits_to_keep: Union[int, torch.Tensor] = 0, |
| vinputs: Optional[torch.Tensor] = None, |
| timestep: Optional[torch.Tensor] = None, |
| token_types: Optional[torch.Tensor] = None, |
| return_mid_results_layers: Optional[list] = None, |
| use_gradient_checkpointing: bool = False, |
| use_gradient_checkpointing_offload: bool = False, |
| **kwargs, |
| ) -> Union[tuple, Qwen3VLCausalLMOutputWithPast]: |
| outputs = self.model( |
| input_ids=input_ids, |
| pixel_values=pixel_values, |
| pixel_values_videos=pixel_values_videos, |
| image_grid_thw=image_grid_thw, |
| video_grid_thw=video_grid_thw, |
| position_ids=position_ids, |
| attention_mask=attention_mask, |
| past_key_values=past_key_values, |
| inputs_embeds=inputs_embeds, |
| cache_position=cache_position, |
| vinputs=vinputs, |
| timestep=timestep, |
| token_types=token_types, |
| return_mid_results_layers=return_mid_results_layers, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| use_gradient_checkpointing_offload=use_gradient_checkpointing_offload, |
| **kwargs, |
| ) |
|
|
| if vinputs is not None: |
| return Qwen3VLCausalLMOutputWithPast( |
| x_pred=outputs.x_pred, |
| mid_results=outputs.mid_results if hasattr(outputs, 'mid_results') else None, |
| ) |
|
|
| hidden_states = outputs[0] |
| slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep |
| logits = self.lm_head(hidden_states[:, slice_indices, :]) |
|
|
| loss = None |
| if labels is not None: |
| loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.text_config.vocab_size) |
|
|
| return Qwen3VLCausalLMOutputWithPast( |
| loss=loss, |
| logits=logits, |
| past_key_values=outputs.past_key_values, |
| rope_deltas=outputs.rope_deltas, |
| ) |
|
|