from __future__ import annotations from dataclasses import dataclass from typing import Any import torch import torch.nn as nn from transformers import PreTrainedModel, Qwen3VLConfig, Qwen3VLTextConfig from transformers.models.qwen3_vl.modeling_qwen3_vl import Qwen3VLModel, Qwen3VLTextModel from transformers.utils import ModelOutput from .configuration_qwen3_vl_stitched import Qwen3VLStitchedConfig from .sequence_adapter import GatedDeltaResidualAdapter, SequenceAdapterConfig @dataclass class StitchedH3ConditioningOutput(ModelOutput): """Raw post-L50 H3 conditioning and optional diagnostic intermediates.""" last_hidden_state: torch.Tensor | None = None source_hidden_state: torch.Tensor | None = None bridge_hidden_state: torch.Tensor | None = None position_ids: torch.Tensor | None = None class Qwen3VLStitchedForH3Conditioning(PreTrainedModel): """8B L1-L16 -> learned bridge -> 32B L30-L50 conditioning frontend.""" config_class = Qwen3VLStitchedConfig base_model_prefix = "" main_input_name = "input_ids" _no_split_modules = [ "Qwen3VLVisionBlock", "Qwen3VLTextDecoderLayer", "GatedDeltaResidualAdapter", "_SequenceLane", ] _supports_sdpa = True _supports_flash_attn = True def __init__(self, config: Qwen3VLStitchedConfig) -> None: super().__init__(config) source_config = Qwen3VLConfig(**config.source_config) source_config.text_config.num_hidden_layers = config.source_layer source_config.text_config.use_cache = False self.source = Qwen3VLModel(source_config) # The bridge consumes raw post-L16 states, before the source model's final norm. self.source.language_model.norm = nn.Identity() bridge_config = SequenceAdapterConfig(**config.bridge_config) self.bridge = GatedDeltaResidualAdapter(bridge_config) target_config = Qwen3VLTextConfig(**config.target_text_config) target_config.num_hidden_layers = config.output_layer - config.target_layer target_config.use_cache = False self.tail = Qwen3VLTextModel(target_config) # Tail input is supplied directly and H3 consumes raw post-L50 states. self.tail.embed_tokens = nn.Identity() self.tail.norm = nn.Identity() def get_input_embeddings(self) -> nn.Module: return self.source.get_input_embeddings() def set_input_embeddings(self, value: nn.Module) -> None: self.source.set_input_embeddings(value) def _position_ids( self, *, input_ids: torch.Tensor | None, inputs_embeds: torch.Tensor | None, attention_mask: torch.Tensor | None, mm_token_type_ids: torch.Tensor | None, image_grid_thw: torch.Tensor | None, video_grid_thw: torch.Tensor | None, ) -> torch.Tensor: if input_ids is not None and mm_token_type_ids is not None: position_ids, _ = self.source.get_rope_index( input_ids=input_ids, mm_token_type_ids=mm_token_type_ids, image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, attention_mask=attention_mask, ) return position_ids reference = input_ids if input_ids is not None else inputs_embeds if reference is None: raise ValueError("input_ids or inputs_embeds are required") batch_size, sequence_length = reference.shape[:2] if attention_mask is None: positions = torch.arange(sequence_length, device=reference.device) positions = positions.view(1, -1).expand(batch_size, -1) else: positions = attention_mask.long().cumsum(dim=-1) - 1 positions = positions.masked_fill(attention_mask == 0, 0) return positions.unsqueeze(0).expand(3, -1, -1) def forward( self, input_ids: torch.LongTensor | None = None, attention_mask: torch.Tensor | None = None, position_ids: torch.LongTensor | None = None, inputs_embeds: torch.FloatTensor | None = None, pixel_values: torch.Tensor | None = None, pixel_values_videos: torch.FloatTensor | None = None, image_grid_thw: torch.LongTensor | None = None, video_grid_thw: torch.LongTensor | None = None, mm_token_type_ids: torch.IntTensor | None = None, use_cache: bool | None = False, output_intermediates: bool = False, return_dict: bool | None = None, **kwargs: Any, ) -> StitchedH3ConditioningOutput | tuple[torch.Tensor, ...]: if use_cache: raise ValueError( "The stitched H3 frontend returns a conditioning sequence and does not " "support KV caching" ) if position_ids is None: position_ids = self._position_ids( input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, mm_token_type_ids=mm_token_type_ids, image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, ) source_output = self.source( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, inputs_embeds=inputs_embeds, pixel_values=pixel_values, pixel_values_videos=pixel_values_videos, image_grid_thw=image_grid_thw, video_grid_thw=video_grid_thw, mm_token_type_ids=mm_token_type_ids, use_cache=False, return_dict=True, **kwargs, ) source_state = source_output.last_hidden_state bridge_state = self.bridge(source_state, attention_mask=attention_mask) tail_output = self.tail( inputs_embeds=bridge_state, attention_mask=attention_mask, position_ids=position_ids, use_cache=False, return_dict=True, **kwargs, ) layer_50 = tail_output.last_hidden_state if return_dict is False: if output_intermediates: return layer_50, source_state, bridge_state, position_ids return (layer_50,) return StitchedH3ConditioningOutput( last_hidden_state=layer_50, source_hidden_state=source_state if output_intermediates else None, bridge_hidden_state=bridge_state if output_intermediates else None, position_ids=position_ids if output_intermediates else None, ) __all__ = ["Qwen3VLStitchedForH3Conditioning", "StitchedH3ConditioningOutput"]