
    ^j                    ~   d dl Z d dlmZ d dlmZ d dlmZ d dlmZ d dl	Z	d dl	m
Z
 d dlmZ dd	lmZ dd
lmZ ddlmZmZ ddlmZ ddlmZmZmZmZmZmZmZm Z  ddl!m"Z" ddl#m$Z$m%Z% ddl&m'Z'm(Z( ddl)m*Z*m+Z+ ddl,m-Z- ddl.m/Z/m0Z0m1Z1m2Z2m3Z3m4Z4 ddl5m6Z6m7Z7m8Z8 ddl9m:Z:m;Z; ddl<m=Z= ddl>m?Z?m@Z@mAZAmBZBmCZCmDZDmEZE ddlFmGZGmHZHmIZImJZJmKZKmLZLmMZMmNZNmOZO ddlPmQZQ ddlRmSZS ddlTmUZU ddlVmWZWmXZXmYZYmZZZ  e2       r	  e3j                  e\      Z]ded e	j                  d!e	j                  dz  d"edz  d#e	j                  dz  d$e	j                  d%e_fd&Z` G d' d(eJ      Za G d) d*eG      Zbe G d+ d,e$             Zce0e G d- d.e%                    Zd G d/ d0e
j                        Zf G d1 d2eM      Zg G d3 d4e
j                        Zh G d5 d6e
j                        Zi G d7 d8e
j                        Zj G d9 d:e
j                        Zk G d; d<e
j                        Zl G d= d>e
j                        Zn G d? d@e
j                        Zo G dA dBe
j                        Zp G dC dDe
j                        Zq G dE dFe
j                        Zr G dG dHeB      Zs	 ddIe	j                  dJe	j                  dKe	j                  d#e	j                  dLetd%e	j                  fdMZu G dN dOeQ      Zve8 G dP dQe?             Zw G dR dSe@      Zx G dT dUe
j                        Zy G dV dWeB      Zz G dX dYeC      Z{ G dZ d[e
j                        Z| G d\ d]eS      Z} G d^ d_e
j                        Z~ G d` dae@      Z G db dceE      Z G dd deeL      Z e0dfg       G dh dieD             Z e0djg       G dk dleA             Z G dm dne      Z G do dpe      Z G dq dreK      Zdse	j                  dz  dte	j                  dz  d%edz  fduZdve	j                  dwe	j                  d%e	j                  fdxZ e0dyg       G dz d{eI             Z e0d|g       G d} d~eH             Zg dZy)    N)UserDict)Callable)	dataclass)cached_property)nn)
functional   )initialization)ACT2FN)CacheDynamicCache)PreTrainedConfig)_preprocess_mask_argumentsblockwise_overlaycreate_bidirectional_maskcreate_causal_maskcreate_masks_for_generate!create_sliding_window_causal_maskmaybe_pad_block_sequence_idssliding_window_overlay)FlashAttentionKwargs)BaseModelOutputWithPastBaseModelOutputWithPooling)ROPE_INIT_FUNCTIONSdynamic_rope_update)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)TransformersKwargsauto_docstringcan_return_tupleis_accelerate_availableloggingtorch_compilable_check)maybe_autocastmerge_with_config_defaultsno_inherit_decorator)OutputRecordercapture_outputs   )	AutoModel)Gemma3AttentionGemma3DecoderLayerGemma3ForCausalLM	Gemma3MLPGemma3RotaryEmbeddingGemma3TextModelGemma3TextScaledWordEmbedding)	Gemma3nCausalLMOutputWithPastGemma3nForConditionalGenerationGemma3nModelGemma3nModelOutputWithPastGemma3nMultimodalEmbedderGemma3nPreTrainedModelGemma3nRMSNormapply_rotary_pos_embeager_attention_forward)LlamaRotaryEmbedding)MixtralExperts)sliding_window_mask_function   )Gemma4AudioConfigGemma4ConfigGemma4TextConfigGemma4VisionConfigconfiginputs_embedsattention_maskpast_key_valuesposition_idsblock_sequence_idsreturnc                     | ||||d}t        di |}t        di |ddi\  }}	}	}	}
}	}|r|}nt        |||
|      }t        di |t        |      t	        | j
                        d}||dS )a  Create full_attention and sliding_attention masks with correct composition.

    For global (full attention) layers:  causal only (no bidirectional)
    For local (sliding window) layers:  AND(sliding_window, OR(causal, blockwise))

    Unlike Gemma 3 (which applies bidirectional attention on all layers), Gemma 4
    explicitly disables bidirectional attention on global attention layers.
    rD   rE   rF   rG   rH   	layer_idxr   )or_mask_functionand_mask_functionfull_attentionsliding_attention )r   r   r   r   r   sliding_window)rD   rE   rF   rG   rH   rI   mask_kwargs	full_mask
early_exit_	kv_length	kv_offsetpadded_block_sequence_idssliding_masks                 t/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/gemma4/modular_gemma4.pycreate_masks_for_vision_modelr^   X   s    " &(*$K #1[1I 4N 4
440J1aAy $6!$@	9%
! & 
*+DE01F1FGL $)     c                   b    e Zd ZU dZdZeeeej                  ej                  f   f   dz  e
d<   y)Gemma4ModelOutputWithPasta  
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
        `past_key_values` input) to speed up sequential decoding.
    image_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        image_hidden_states of the model produced by the vision encoder and after projecting the last hidden state.
    audio_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        audio_hidden_states of the model produced by the audio encoder and after projecting the last hidden state.
    shared_kv_states (`dict`, *optional*):
        Dictionary mapping layer type strings to tuples of (key_states, value_states) tensors.
        Used to pass shared KV states between layers during KV sharing.
    Nshared_kv_states__name__
__module____qualname____doc__rb   dictstrtupletorchTensor__annotations__rS   r_   r]   ra   ra      s7    " MQd3ellELL&@ AABTIPr_   ra   c                   b    e Zd ZU dZdZeeeej                  ej                  f   f   dz  e
d<   y)Gemma4CausalLMOutputWithPasta  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction).
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.text_config.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
        `past_key_values` input) to speed up sequential decoding.
    image_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        image_hidden_states of the model produced by the vision encoder after projecting last hidden state.
    audio_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        audio_hidden_states of the model produced by the audio encoder and after projecting the last hidden state.
    shared_kv_states (`dict`, *optional*):
        Dictionary mapping layer type strings to tuples of (key_states, value_states) tensors.
        Used to pass shared KV states between layers during KV sharing.
    Nrb   rc   rS   r_   r]   ro   ro      s7    * MQd3ellELL&@ AABTIPr_   ro   c                   b    e Zd ZU dZdZeeeej                  ej                  f   f   dz  e
d<   y)Gemma4TextModelOutputWithPasta9  
    BaseModelOutputWithPast extended with shared_kv_states for KV sharing.

    Args:
        shared_kv_states (`dict`, *optional*):
            Dictionary mapping layer type strings to tuples of (key_states, value_states) tensors.
            Used to pass shared KV states between layers during KV sharing.
    Nrb   rc   rS   r_   r]   rq   rq      s7     MQd3ellELL&@ AABTIPr_   rq   c                   :    e Zd ZU dZdZej                  dz  ed<   y)Gemma4AudioModelOutputz
    attention_mask (`torch.BoolTensor`, *optional*):
        A torch.BoolTensor of shape `(batch_size, num_frames)`. True for valid positions, False for padding.
    NrF   )rd   re   rf   rg   rF   rk   
BoolTensorrm   rS   r_   r]   rs   rs      s    
 /3NE$$t+2r_   rs   c                   n     e Zd Zdeez  dededdf fdZdej                  dej                  fdZ	 xZ
S )	Gemma4ClippableLinearrD   in_featuresout_featuresrJ   Nc                    t         |           |j                  | _        t        j                  ||d      | _        | j                  r| j                  dt        j                  t        d                    | j                  dt        j                  t        d                   | j                  dt        j                  t        d                    | j                  dt        j                  t        d                   y y )NFbias	input_mininf	input_max
output_min
output_max)
super__init__use_clipped_linearsr   Linearlinearregister_bufferrk   tensorfloat)selfrD   rw   rx   	__class__s       r]   r   zGemma4ClippableLinear.__init__   s     	#)#=#= ii\F##  ellE%L=.IJ  ell5<.HI  u||U5\M/JK  u||E%L/IJ	 $r_   hidden_statesc                    | j                   r+t        j                  || j                  | j                        }| j                  |      }| j                   r+t        j                  || j                  | j                        }|S N)r   rk   clampr|   r~   r   r   r   )r   r   s     r]   forwardzGemma4ClippableLinear.forward   s\    ##!KKt~~t~~VMM2##!KKtXMr_   )rd   re   rf   rC   r@   intr   rk   rl   r   __classcell__r   s   @r]   rv   rv      sT    K"%66K K 	K
 
K 	U\\ 	ell 	r_   rv   c                       e Zd Zy)Gemma4RMSNormNrd   re   rf   rS   r_   r]   r   r          r_   r   c                        e Zd ZU dZej
                  ed<   def fdZ ej                         dej
                  dej
                  fd       Z
 xZS ) Gemma4AudioRelPositionalEncodingzSinusoidal relative positional encoding for the audio encoder.

    Produces position embeddings of shape [1, context_size // 2 + 1, hidden_size] with
    concatenated [sin..., cos...] layout matching the original Gemma4 convention.
    inv_timescalesrD   c                    t         |           |j                  | _        |j                  |j                  z   dz
  |j
                  z   | _        d}d}| j                  dz  }t        j                  ||z        t        |dz
  d      z  }|t        j                  t        j                  |      | z        z  }| j                  d|j                  d      j                  d      d       y )	Nr?         ?     @r*   r   r   F
persistent)r   r   hidden_sizeattention_chunk_sizeattention_context_leftattention_context_rightcontext_sizemathlogmaxrk   exparanger   	unsqueeze)r   rD   min_timescalemax_timescalenum_timescaleslog_timescale_incrementr   r   s          r]   r   z)Gemma4AudioRelPositionalEncoding.__init__   s    !--''&*G*GG!KfNlNll 	 ))Q."&((==+H"ICP^abPbdeLf"f&5<<3OSjRj3j)kk-~/G/G/J/T/TUV/Wdijr_   r   rJ   c                 t   t        j                  | j                  dz  dd|j                        }|d   }|| j                  j                  |j                        z  }t        j                  t        j                  |      t        j                  |      gd      }|j                  |j                        S )Nr*   device.Ndimdtype)
rk   r   r   r   r   tocatsincosr   )r   r   rH   scaled_time	pos_embeds        r]   r   z(Gemma4AudioRelPositionalEncoding.forward  s    ||D$5$5$:B=K_K_`#I."T%8%8%;%;=CWCW%;%XXIIuyy5uyy7MNTVW	||-"5"5|66r_   )rd   re   rf   rg   rk   rl   rm   r@   r   no_gradr   r   r   s   @r]   r   r      sU     LL k0 k U]]_7U\\ 7ell 7 7r_   r   c                   P    e Zd ZdZdedef fdZdej                  dej                  fdZ	dej                  dej                  fdZ
d	ej                  dej                  fd
Z	 ddej                  dej                  dej                  dz  deej                  df   fdZ xZS )Gemma4AudioAttentionz3Chunked local attention with relative position biasrD   rM   c                     t         |           || _        || _        |j                  | _        |j                  |j                  z  | _        |j                  | _	        | j                  dz  t        j                  d      z  | _        t        j                  dt        j                  z         t        j                  d      z  | _        |j                  | _        |j"                  dz
  | _        |j&                  | _        | j                   | j$                  z   | j(                  z   | _        t-        ||j                  | j                  | j                  z        | _        t-        ||j                  | j                  | j                  z        | _        t-        ||j                  | j                  | j                  z        | _        t-        ||j                  |j                        | _        t7        j8                  |j                  | j                  | j                  z  d      | _        t7        j<                  t?        j@                  | j                              | _!        | jE                  dt?        jF                  | j
                        d       y )N      r*   r?   Frz   softcapr   )$r   r   rD   rM   attention_logit_capattention_logits_soft_capr   num_attention_headshead_dim	num_headsr   r   q_scaleek_scaler   
chunk_sizer   max_past_horizonr   max_future_horizonr   rv   q_projk_projv_projpostr   r   relative_k_proj	Parameterrk   zerosper_dim_scaler   r   r   rD   rM   r   s      r]   r   zGemma4AudioAttention.__init__  s   ")/)C)C&**f.H.HH33t+txx{:xxDFF
+dhhqk9 55 & = = A"("@"@ OOd.C.CCdF]F]]+FF4F4FY]YfYfHfg+FF4F4FY]YfYfHfg+FF4F4FY]YfYfHfg)&&2D2DfFXFXY	!yy););T^^dmm=[bgh\\%++dmm*DEYT5S5S(Tafgr_   r   rJ   c           	         |j                   \  }}}}|| j                  z   dz
  | j                  z  }|| j                  z  |z
  }t        j                  |ddddd|f      }|j	                  ||| j                  ||      j                         S )zSplits a `(batch_size, seq_len, num_heads, head_dim)` tensor into non-overlapping blocks of `chunk_size` along the sequence dim.r?   r   )shaper   Fpadreshape
contiguous)r   r   
batch_sizeseq_lenr   r   
num_blocksr   s           r]   _convert_to_blockz&Gemma4AudioAttention._convert_to_block3  s    3@3F3F0
GY/!3G
4??*W4maAq!S-AB$$ZT__iYabmmoor_   c           
      @   |j                   \  }}}}t        j                  |dddd| j                  | j                  | j
                  z   dz
  f      }|j                  d| j                  | j
                        }t        j                  |dd      }|j                         S )z`Extracts overlapping context windows of `context_size` for every block, strided by `chunk_size`.r   r?   r   r*   )r   r   r   r   r   r   unfoldr   rk   movedimr   )r   r   r   r   r   r   s         r]   _extract_block_contextz+Gemma4AudioAttention._extract_block_context;  s    3@3F3F0
GYAq!Q(=(=t?V?VY]YhYh?hkl?lm
 &,,Q0A0A4??SmR;''))r_   xc                     |j                   \  }}}}}| j                  }t        j                  |d|dz   |z
  f      }|j	                  |||||dz   z        }|dd||z  f   }|j	                  |||||      S )zjRelative position shift for blocked attention. See appendix B of https://huggingface.co/papers/1901.02860.r   r?   .N)r   r   r   r   view)r   r   r   r   r   
block_sizeposition_lengthr   s           r]   
_rel_shiftzGemma4AudioAttention._rel_shiftE  s    IJF
Iz:((EE!a)O;<=FF:y*jLSTDT6UVc.Z,.../vvj)Z\RRr_   Nposition_embeddingsrF   c                    |j                   \  }}}||| j                  | j                  f}| j                  |      j	                         j                  |      }| j                  |      j	                         j                  |      }	| j                  |      j	                         j                  |      }
|| j                  z  t        j                  | j                        z  }|	| j                  z  }	| j                  |      }| j                  |	      }	| j                  |
      }
|j                   d   }| j                  |      }|j                  d| j                  | j                        }|j!                  |j"                        }|j%                  ddddd      }||	j%                  ddddd      z  }|j'                  || j                  d| j                        }||j%                  ddd      z  }|j'                  || j                  || j(                  d      }| j+                  |      }||z   }|| j,                  z  }t/        j0                  |      }|| j,                  z  }|4|j3                  |j5                         | j6                  j8                        }t        j:                  |dt.        j<                        j!                  |
j"                        }||
j%                  ddddd      z  }|j%                  ddddd      j'                  ||| j(                  z  d      }|d d d |f   j?                         }| jA                  |j!                  |j"                              }||fS )	Nr?   r   r   r   r	   r*      r   r   )!r   r   r   r   r   r   r   r   r   r   softplusr   r   r   r   r   r   r   permuter   r   r   r   rk   tanhmasked_filllogical_notrD   attention_invalid_logits_valuesoftmaxfloat32r   r   )r   r   r   rF   r   
seq_lengthrX   hidden_shapequery_states
key_statesvalue_statesr   relative_key_statesqueries	matrix_acqueries_flat	matrix_bdattn_weightsattn_outputs                      r]   r   zGemma4AudioAttention.forwardN  s    %2$7$7!
J"JN{{=1779>>|L[[/557<<\J
{{=1779>>|L#dll2QZZ@R@R5SS$,,.
--l;00<
22<@!''*
"223FG166r4>>4==Y144<;M;M4N&&q!Q15j00Aq!Q??	z4>>2t}}U #6#>#>q!Q#GG	%%j$..*doo_ab	OOI.	 9,#dll2zz,/#dll2%'33**,dkk.X.XL yy2U]]KNN|OaOab"\%9%9!Q1a%HH!))!Q1a8@@ZZ^ZiZiMikmn!![j[.1<<>ii}/B/B CDL((r_   r   )rd   re   rf   rg   r@   r   r   rk   rl   r   r   r   rt   rj   r   r   r   s   @r]   r   r     s    =h0 hS h4pu|| p p*ELL *U\\ *SELL SU\\ S 37	1)||1) #\\1) ((4/	1)
 
u||T!	"1)r_   r   c                   ^     e Zd Z fdZddej
                  dej
                  dz  fdZ xZS )'Gemma4AudioSubSampleConvProjectionLayerc                     t         |           t        j                  ||dddd      | _        t        j
                  ||dd      | _        t        j                         | _        y )N)r	   r	   )r*   r*   r?   F)in_channelsout_channelskernel_sizestridepaddingr{   T)epselementwise_affiner{   )	r   r   r   Conv2dconv	LayerNormnormReLUact)r   r  r  norm_epsr   s       r]   r   z0Gemma4AudioSubSampleConvProjectionLayer.__init__  sW    II#%
	 LL8PT[`a	779r_   Nr   maskc           
         |,|j                  |j                        }||d d d d d d f   z  }| j                  |j                  | j                  j                  j                              }| j                  | j                  |j                  dddd            j                  dddd      j                               }||d d d d df   }||fS )Nr   r   r*   r	   r?   )	r   r   r  weightr   r  r  r   r   )r   r   r  s      r]   r   z/Gemma4AudioSubSampleConvProjectionLayer.forward  s    77-"6"677D)DD!T1A,BBM		-"2"24993C3C3I3I"JK=+@+@Aq!+L!M!U!UVWYZ\]_`!a!l!l!no3Q3<Dd""r_   r   )rd   re   rf   r   rk   rl   r   r   r   s   @r]   r  r    s(    #U\\ #9L #r_   r  c            	            e Zd Zdef fdZ	 ddej                  dej                  dz  deej                  ej                  f   fdZ xZ	S )	"Gemma4AudioSubSampleConvProjectionrD   c                 v   t         |           t        d|j                  d   |j                        | _        t        |j                  d   |j                  d   |j                        | _        |j                  d   dz  |j                  d   z  }t        j                  ||j                  d      | _
        y )Nr?   r   )r  r  r  r   Frz   )r   r   r  subsampling_conv_channelsrms_norm_epslayer0layer1r   r   r   input_proj_linear)r   rD   proj_input_dimr   s      r]   r   z+Gemma4AudioSubSampleConvProjection.__init__  s    =99!<((

 >88;99!<((

 !::1=BfFfFfghFii!#>6;M;MTY!Zr_   Ninput_featuresinput_features_maskrJ   c                 &   |j                  d      }| j                  ||      \  }}| j                  ||      \  }}|j                  \  }}}}|j	                  dddd      j                         j                  ||d      }| j                  |      |fS )Nr?   r   r*   r	   r   )r   r  r  r   r   r   r   r  )r   r   r!  r   r  r   rX   r   s           r]   r   z*Gemma4AudioSubSampleConvProjection.forward  s    
 '003"kk-9LMt"kk->t$1$7$7!
Aw%--aAq9DDFNNz[bdfg%%m4d::r_   r   )
rd   re   rf   r@   r   rk   rl   rj   r   r   r   s   @r]   r  r    sW    [0 [$ 48;; #\\D0; 
u||U\\)	*	;r_   r  c                   \     e Zd Zdef fdZdej                  dej                  fdZ xZS )Gemma4AudioFeedForwardrD   c                    t         |           || _        t        ||j                  |j                  dz        | _        t        ||j                  dz  |j                        | _        t        |j                        | _        t        |j                        | _	        t        |j                     | _        |j                  | _        |j                  | _        y )Nr   )r   r   rD   rv   r   ffw_layer_1ffw_layer_2r   pre_layer_normpost_layer_normr   
hidden_actact_fngradient_clippingresidual_weightpost_layer_scaler   rD   r   s     r]   r   zGemma4AudioFeedForward.__init__  s    09K9KVM_M_bcMcd09K9Ka9OQWQcQcd+F,>,>?,V-?-?@V../!'!9!9 & 6 6r_   r   rJ   c                    t        | j                  t        j                  |j                        j
                        }|}t        j                  || |      }| j                  |      }| j                  |      }| j                  |      }| j                  |      }t        j                  || |      }| j                  |      }|| j                  z  }||z  }|S r   )minr,  rk   finfor   r   r   r(  r&  r+  r'  r)  r.  )r   r   r,  residuals       r]   r   zGemma4AudioFeedForward.forward  s     6 6MDWDW8X8\8\] M4E3EGXY++M:((7M2((7M4E3EGXY,,];...!r_   	rd   re   rf   r@   r   rk   rl   r   r   r   s   @r]   r$  r$    s+    70 7U\\ ell r_   r$  c                   `     e Zd Zed        Zdej                  dej                  f fdZ xZS )Gemma4AudioCausalConv1dc                 p    | j                   d   dz
  | j                  d   z  dz   }|| j                  d   z
  S )Nr   r?   )r  dilationr	  )r   effective_kernel_sizes     r]   left_padz Gemma4AudioCausalConv1d.left_pad  s>    !%!1!1!!4q!8DMM!<L Lq P$t{{1~55r_   r   rJ   c                 z    t         j                  j                  || j                  df      }t        |   |      S )Nr   )r   r   r   r:  r   r   )r   r   r   s     r]   r   zGemma4AudioCausalConv1d.forward  s3     MMa$--!34wq!!r_   )	rd   re   rf   r   r:  rk   rl   r   r   r   s   @r]   r6  r6    s;     6 6"<<" 
	" "r_   r6  c                   \     e Zd Zdef fdZdej                  dej                  fdZ xZS )Gemma4AudioLightConv1drD   c                 6   t         |           || _        t        ||j                  |j                  dz        | _        t        ||j                  |j                        | _        t        |j                  |j                  |j                  |j                  d      | _	        t        |j                  |j                  d      | _        t        |j                  |j                  d      | _        t        |j                     | _        |j"                  | _        y )Nr*   F)r  r  r  groupsr{   Tr  
with_scale)r   r   rD   rv   r   linear_start
linear_endr6  conv_kernel_sizedepthwise_conv1dr   r  r(  	conv_normr   r*  r+  r,  r/  s     r]   r   zGemma4AudioLightConv1d.__init__  s    1&&:L:LfN`N`cdNde/8J8JFL^L^_ 7**++//%%!
 ,F,>,>FDWDWdhi&v'9'9v?R?R_cdV../!'!9!9r_   r   rJ   c                    |}| j                  |      }| j                  |      }t        j                  j	                  |d      }| j                  |j                  dd            j                  dd      }t        | j                  t        j                  |j                        j                        }t        j                  || |      }| j                  |      }| j                  |      }| j!                  |      }||z  }|S )Nr   r   r?   r*   )r(  rB  r   r   glurE  	transposer1  r,  rk   r2  r   r   r   rF  r+  rC  )r   r   r3  r,  s       r]   r   zGemma4AudioLightConv1d.forward  s     ++M:))-8))-R)@--m.E.Ea.KLVVWXZ[\   6 6MDWDW8X8\8\]M4E3EGXY}5M26!r_   r4  r   s   @r]   r=  r=    s+    :0 :(U\\ ell r_   r=  c            
            e Zd Zdedef fdZdej                  dej                  dz  dej                  de	e
   d	ej                  f
d
Z xZS )Gemma4AudioLayerrD   rM   c                 p   t         |           || _        t        |      | _        t        |      | _        t        ||      | _        t        |      | _	        t        |j                        | _        t        |j                        | _        t        |j                        | _        |j                  | _        y r   )r   r   rD   r$  feed_forward1feed_forward2r   	self_attnr=  lconv1dr   r   norm_pre_attnnorm_post_attnnorm_outr,  r   s      r]   r   zGemma4AudioLayer.__init__+  s    3F;3F;-fi@-f5*6+=+=>+F,>,>?%f&8&89!'!9!9r_   r   rF   Nr   kwargsrJ   c                 @   t        | j                  t        j                  | j                  j
                  j                        j                        }| j                  |      }|}t        j                  || |      }| j	                  |      }| j                  |||      \  }}t        j                  || |      }| j                  |      }||z  }| j                  |      }| j                  |      }t        j                  || |      }| j                  |      }|S )N)r   r   rF   )r1  r,  rk   r2  rQ  r  r   r   rM  r   rO  rR  rP  rN  rS  )r   r   rF   r   rT  r,  r3  rX   s           r]   r   zGemma4AudioLayer.forward:  s      6 6DDVDVD]D]DcDc8d8h8hi**=9 M4E3EGXY**=9>>' 3) * 
q M4E3EGXY++M:!]3**=9M4E3EGXYm4r_   )rd   re   rf   r@   r   r   rk   rl   rt   r   r   r   r   r   s   @r]   rK  rK  *  si    :0 :S : ||  ((4/  #\\	 
 +,  
 r_   rK  c                        e Zd Zdef fdZdej                  dej                  dej                  fdZdej                  dej                  dej                  dej                  fdZ xZ	S )	Gemma4VisionPatchEmbedderrD   c                    t         |           || _        |j                  | _        |j                  | _        |j
                  | _        t        j                  d| j                  dz  z  | j                  d      | _        t        j                  t        j                  d| j
                  | j                              | _        y )Nr	   r*   Frz   )r   r   rD   r   
patch_sizeposition_embedding_sizer   r   
input_projr   rk   onesposition_embedding_tabler/  s     r]   r   z"Gemma4VisionPatchEmbedder.__init__a  s    !-- ++'-'E'E$))A(:$:D<L<LSXY(*UZZ4C_C_aeaqaq5r(s%r_   pixel_position_idspadding_positionsrJ   c                    |j                  d      }t        j                  |d   | j                  d         }t        j                  |d   | j                  d         }||z   }t	        j
                  |j                  d      d|      }|S )ak  Compute 2-D patch position embeddings via embedding lookup.

        ``pixel_position_ids`` has shape ``(batch, num_patches, 2)`` where the
        last dimension holds (x, y) indices into ``position_embedding_table``
        (shape ``(2, position_embedding_size, hidden_size)``).  The result is the
        sum of the x- and y-embeddings for each patch.
        r   r1  .r   .r?   r?   r           )r   r   	embeddingr]  rk   wherer   )r   r^  r_  clamped_positionsx_emby_embr   s          r]   _position_embeddingsz.Gemma4VisionPatchEmbedder._position_embeddingsk  s     /444; -f5t7T7TUV7WX-f5t7T7TUV7WX#em#kk*;*E*Eb*I3Pcd""r_   pixel_valuesc                     d|dz
  z  }| j                   j                  j                  x}j                  r|j	                  |      }| j                  |      }| j                  ||      }||z   S )Nr*         ?)r[  r  r   is_floating_pointr   rj  )r   rk  r^  r_  target_dtyper   r   s          r]   r   z!Gemma4VisionPatchEmbedder.forward  sn     L3./ OO22888LKK'??<8L5"778JL]^222r_   )
rd   re   rf   rC   r   rk   rl   rj  r   r   r   s   @r]   rW  rW  `  su    t1 t#u|| #X]XdXd #iniuiu #,	3!LL	3>Cll	3_d_k_k	3		3r_   rW  c                   .    e Zd ZdZdef fdZdej                  dej                  dede	ej                  ej                  f   fdZ
	 ddej                  dej                  d
ej                  ded	z  de	ej                  ej                  f   f
dZ xZS )Gemma4VisionPoolera[  Spatial pooling and ``sqrt(hidden_size)`` scaling for vision encodings.

    The scaling expands the activation magnitude, which can exceed the float16 range, so it is
    computed in float32 and the pooled features are returned in float32. The caller
    (``Gemma4VisionModel.forward``) standardizes them and casts back to the working dtype.
    rD   c                 l    t         |           |j                  | _        | j                  dz  | _        y )Nrm  )r   r   r   root_hidden_sizer/  s     r]   r   zGemma4VisionPooler.__init__  s/    !-- $ 0 0# 5r_   r   r^  lengthrJ   c                    |j                   d   }t        ||z  dz        }|dz  }||z  |k7  r%t        d|j                    d| d|d|d| d	      |j                  d
      }|d   j	                  dd      d
   dz   }t        j                  ||d      }	|	d   ||z  |	d   z  z   }	t        j                  |	j                         |      j                         |z  }
|
j                  dd      |j                         z  }t        j                  |
d
k(  j                  d            }|j                  |j                        |fS )z
        2D spatial pooling according to patch positions.
        Pools the input tokens by averaging patches within a `k^2` grid, where `k` is determined by the ratio between
        input and output lengths
        r?   rm  r*   zCannot pool z to z: k=z^2 times length=z	 must be .r   ra  rb  r   Tr   keepdimfloor)rounding_moderc  r   )r   r   
ValueErrorr   r   rk   divr   one_hotlongr   rI  r   allr   r   )r   r   r^  rt  input_seq_lenk	k_squaredrg  max_xkernel_idxsweightsoutputr  s                r]   _avg_pool_by_positionsz)Gemma4VisionPooler._avg_pool_by_positions  sh    &++A.&(S01qD	v.}2234xu!EVviW`an`oopq  /444;!&)--"d-CAFJii 11GL!&)UaZ;v;N,NN))K,,.7==?)K""1a(=+>+>+@@  'Q,!3!3!3!:;yy,,-t33r_   Nr_  output_lengthc                 8   ||j                   d   kD  rt        d| d|j                   d    d      |j                  |j                  d      d      }|j                   d   |k7  r| j	                  |||      \  }}|j                         | j                  z  }||fS )Nr?   z*Cannot output more soft tokens (requested z) than there are patches (z9). Change the value of `num_soft_tokens` when processing.r   rd  )r   r{  r   r   r  r   rs  )r   r   r^  r_  r  s        r]   r   zGemma4VisionPooler.forward  s     =..q11<]O L"((+,,eg 
 &112C2M2Mb2QSVWq!]2/3/J/J1=0,M, &++-0E0EE///r_   r   )rd   re   rf   rg   rC   r   rk   rl   r   rj   r  r   r   r   s   @r]   rq  rq    s    61 6
4"\\4?D||4UX4	u||U\\)	*4@ %)0||0 "LL0 !<<	0
 Tz0 
u||U\\)	*0r_   rq  c                   $     e Zd Zdef fdZ xZS )Gemma4VisionMLPrD   c                 
   t         |   | |       t        || j                  | j                        | _        t        || j                  | j                        | _        t        || j                  | j                        | _        y r   )r   r   rv   r   intermediate_size	gate_projup_proj	down_projr/  s     r]   r   zGemma4VisionMLP.__init__  sf    v&.vt7G7GI_I_`,VT5E5EtG]G]^.vt7M7MtO_O_`r_   )rd   re   rf   rC   r   r   r   s   @r]   r  r    s    a1 a ar_   r  r   r   r   unsqueeze_dimc           	         |j                   d   }| j                   d   }d|d|z  z  z  }|dk  rt        d| d| d| d      |g|z  }t        j                  | |d      }	t        j                  ||d      }
t        j                  ||d      }t	        |      D cg c]  }t        |	|   |
|   ||   |	       }}t        j                  |d      S c c}w )
ak  Applies multidimensional RoPE to inputs.

    Args:
        x (`torch.Tensor`): The tensor to embed.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        position_ids (`torch.Tensor`, *optional*):
            If position_ids.ndim + 2 == x.ndim, then this function passes through to `apply_rotary_pos_emb()`.
            Otherwise, position_ids is used to split the inputs, x, into multiple pieces, where each piece is fed to
            `apply_rotary_pos_emb()`, and then concatenated back together.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.

    Returns:
      Tensor of shape [B, L, N, H] with RoPE applied.
    r   r*   r   zEInvalid configuration: num_rotated_channels_per_dim must be > 0, got z (num_input_channels=z, ndim=)r   )r   r   r   r  )r   r{  rk   splitranger:   r   )r   r   r   rH   r  ndimnum_input_channelsnum_rotated_channels_per_dimsplit_sizesx_parts	cos_parts	sin_partsr  y_partss                 r]   apply_multidimensional_roper    s   8 b!D#$(:q4x(H#I #q(,--BCUBV WF!
 	
 0047Kkk![b1GC"5IC"5I t  	aj!!'		
G  99W"%%s   Cc                       e Zd Ze	 	 	 d	dedz  dej                  dz  dedz  dede	f   fd       Z
 ej                         ed               Zy)
Gemma4VisionRotaryEmbeddingNrD   r   r   rJ   ztorch.Tensorc                 $   | j                   d   }t        | dd      xs | j                  | j                  z  }|dz  }d}d|t	        j
                  d|dt        j                        j                  |t        j                        |z  z  z  }||fS )	a  
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            device (`torch.device`):
                The device to use for initialization of the inverse frequencies.
            seq_len (`int`, *optional*):
                The current sequence length. Unused for this type of RoPE.
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        
rope_thetar   Nr*   r   r   r   )r   r   )	rope_parametersgetattrr   r   rk   r   int64r   r   )rD   r   r   baser   spatial_dimattention_factorinv_freqs           r]   compute_default_rope_parametersz;Gemma4VisionRotaryEmbedding.compute_default_rope_parameters  s    & %%l3fj$/c63E3EIcIc3c QhQQekkBEEV[`[f[fEgjuuw
 )))r_   c                 .   | j                   d d d d f   j                         j                  |j                  d   dd      j	                  |j
                        }t        |j
                  j                  t              r/|j
                  j                  dk7  r|j
                  j                  nd}g g }}t        d      D ]  }|d d d d |f   }|d d d d d f   j                         }	t        |d      5  |j                         |	j                         z  j                  dd      }
t        j                  |
|
fd	      }|j                         | j                  z  }|j!                         | j                  z  }d d d        |j#                         |j#                          t        j                  |d	      j	                  |j$                  
      }t        j                  |d	      j	                  |j$                  
      }||fS # 1 sw Y   xY w)Nr   r   r?   mpscpur*   F)device_typeenabledr   r   )r  r   expandr   r   r   
isinstancetyperi   r  r%   rI  rk   r   r   attention_scalingr   appendr   )r   r   rH   inv_freq_expandedr  all_cosall_sinidim_position_idsdim_position_ids_expandedfreqsembr   r   s                 r]   r   z#Gemma4VisionRotaryEmbedding.forward4  s    !MM$4-8>>@GGHZHZ[\H]_acdehhijiqiqr'1!((--'E!((--[`J`ahhmmfk rq 
	 A+Aq!G4(8D!(D(J(J(L%KG 9*0025N5T5T5VVaabcefgiiB7ggi$"8"88ggi$"8"88	9
 NN3NN3
	  iiR(++!''+:iiR(++!''+:Cx9 9s   4BHH	NNN)rd   re   rf   staticmethodrC   rk   r   r   rj   r   r  r   r   r   rS   r_   r]   r  r    s    ,0&*" *"T) *t# * t * 
~u$	%	 *  *D U]]_  r_   r  c                       e Zd Zdedef fdZ	 	 	 ddej                  dej                  dej                  dz  dej                  dz  d	e	e
   d
eej                  ej                  dz  eej                     dz  f   fdZ xZS )Gemma4VisionAttentionrD   rM   c                 6   t         |   | ||       | `| `| `d| _        d| _        t        ||j                  |j                  | j                  z        | _        t        ||j                  |j                  | j                  z        | _        t        ||j                  |j                  | j                  z        | _        t        ||j                  | j                  z  |j                        | _        t!        | j                  |j"                  d      | _        y )Nr   Fr@  )r   r   attn_logit_softcappingrT   
is_slidingscaling	is_causalrv   r   num_key_value_headsr   r   r   r   r   o_projr   r  v_normr   s      r]   r   zGemma4VisionAttention.__init__O  s    vy1'O+FF4F4FHbHbeiererHrs+FF4F4FHbHbeiererHrs+FF4F4FHbHbeiererHrs+FF4N4NQUQ^Q^4^`f`r`rs#DMMv7J7JW\]r_   Nr   r   rF   rH   rT  rJ   c                 N   |j                   d d }g |d| j                  }|\  }}	| j                  |      j                  |      }
| j	                  |
      }
t        |
||	|      }
|
j                  dd      }
| j                  |      j                  |      }| j                  |      }t        |||	|      }|j                  dd      }| j                  |      j                  |      }| j                  |      }|j                  dd      }t        j                  | j                  j                  t              } || |
|||f| j                   r| j"                  nd| j$                  d|\  }} |j&                  g |d j)                         }| j+                  |      }||fS )Nr   r?   r*   rd  )dropoutr  )r   r   r   r   q_normr  rI  r   k_normr   r  r   get_interfacerD   _attn_implementationr;   trainingattention_dropoutr  r   r   r  )r   r   r   rF   rH   rT  input_shaper   r   r   r   r   r   attention_interfacer  r  s                   r]   r   zGemma4VisionAttention.forward\  s    $))#2.88b8$--8&S{{=166|D{{<02<c<X#--a3[[/44\B
[[,
0S#|T
))!Q/
{{=166|D{{<0#--a3(?(M(MKK,,.E)
 %8	%
 /3mmD**LL	%
 	%
!\ *k));;;;FFHkk+.L((r_   r  )rd   re   rf   rC   r   r   rk   rl   
LongTensorr   r   rj   r   r   r   s   @r]   r  r  M  s    ^1 ^c ^  -1.204,)||,) #\\,) t+	,)
 &&-,) +,,) 
u||U\\D0%2E2LL	M,)r_   r  c                       e Zd Zdedef fdZ	 	 	 ddej                  dej                  dej                  dz  dej                  dz  d	e	e
   d
eej                  eej                  ej                  f   dz  f   fdZ xZS )Gemma4VisionEncoderLayerrD   rM   c                 l    t         |   | ||       t        ||      | _        t	        |      | _        y NrD   rM   )r   r   r  rO  r  mlpr   s      r]   r   z!Gemma4VisionEncoderLayer.__init__  s.    vy1.f	R"6*r_   Nr   r   rF   rH   rT  rJ   c                     |}| j                  |      } | j                  d||||d|\  }}| j                  |      }||z   }|}| j                  |      }| j	                  |      }| j                  |      }||z   }|S )N)r   r   rF   rH   rS   )input_layernormrO  post_attention_layernormpre_feedforward_layernormr  post_feedforward_layernorm)r   r   r   rF   rH   rT  r3  rX   s           r]   r   z Gemma4VisionEncoderLayer.forward  s     !,,];)4>> 
' 3)%	

 
q 55mD =0 66}E/77F =0r_   r  )rd   re   rf   rC   r   r   rk   rl   r  r   r   rj   FloatTensorr   r   r   s   @r]   r  r    s    +1 +c + -1.204|| #\\ t+	
 &&- +, 
u  %(9(95;L;L(L"MPT"TT	Ur_   r  c                        e Zd Zdef fdZ	 d
dej                  dej                  dej                  dz  dee	   de
f
d	Z xZS )Gemma4VisionEncoderrD   c           	         t         |           || _        |j                  | _        t        |      | _        t        j                  t        | j                        D cg c]  }t        ||       c}      | _        y c c}w r  )r   r   rD   num_hidden_layers
num_layersr  
rotary_embr   
ModuleListr  r  layers)r   rD   r  r   s      r]   r   zGemma4VisionEncoder.__init__  sc     225f=mmKPQUQ`Q`Kaba%VqAb
bs   A?NrE   rF   r^  rT  rJ   c                     t        | j                  ||      }|}| j                  ||      }| j                  d| j                  j                   D ]  } ||f|||d|} t        |      S )z
        pixel_position_ids (torch.Tensor):
            Patch positions as (x, y) coordinates in the image as [batch, num_patches, 2].
        )rD   rE   rF   N)rF   r   rH   last_hidden_state)r   rD   r  r  r  r   )r   rE   rF   r^  rT  r   r   decoder_layers           r]   r   zGemma4VisionEncoder.forward  s     3;;')
 &"oom=OP "[[)H4;;+H+HI 	M)-$7/	
 M	 'GGr_   r   )rd   re   rf   rC   r   rk   rl   r  r   r   r   r   r   r   s   @r]   r  r    si    
1 
 7;	H||H H ",,t3	H
 +,H 
!Hr_   r  c                   (     e Zd Zdedef fdZ xZS )Gemma4TextMLPrD   rM   c                     |j                   |j                  z
  }||cxk\  xr dkD  nc }|j                  xr |}t        |           |j
                  |r
dz  | _        y dz  | _        y )Nr   r*   r?   )r  num_kv_shared_layersuse_double_wide_mlpr   r   r  )r   rD   rM   first_kv_shared_layer_idxis_kv_shared_layerr  r   s         r]   r   zGemma4TextMLP.__init__  si    $*$<$<v?Z?Z$Z!&*CGaG$88O=O!'!9!9BUQ!][\!]r_   )rd   re   rf   rB   r   r   r   r   s   @r]   r  r    s     ^/ ^C ^ ^r_   r  c                       e Zd ZddefdZy)Gemma4TextRotaryEmbeddingNrD   c                    t         j                  j                  |        |j                  | _        |j                  | _        || _        t        |j                        | _        i | _	        i | _
        | j                  D ]  }| j                  j                  |   }||d   x}dk7  r
t        |   }n| j                  }|| j                  |<   || j                  |<   ||d}|dk(  r
|dk(  rd|d<    || j                  fi |\  }}	| j                  | d|d	
       | j                  | d|j                         d	
       t!        | | d|	        y )N	rope_typedefault)r   
layer_typerQ   proportionalglobal_head_dimhead_dim_key	_inv_freqFr   _original_inv_freq_attention_scaling)r   Moduler   max_position_embeddingsmax_seq_len_cachedoriginal_max_seq_lenrD   setlayer_typesrope_init_fnsr  r  r   r  r   clonesetattr)
r   rD   r   r  rope_paramsr  rope_init_fnrope_init_fn_kwargscurr_inv_freqcurr_attention_scalings
             r]   r   z"Gemma4TextRotaryEmbedding.__init__  sk   
		4 "("@"@$*$B$B!v112SU)+** 	UJ++55jAK"(55	)C29=#CC-9Dz*)2DNN:&-3:"N--)~2M6G#N34@4dPc4d1M1  J<y!9=UZ [  J</A!BMDWDWDYfk lDZL(:;=ST)	Ur_   NN)rd   re   rf   rB   r   rS   r_   r]   r  r    s    U/ Ur_   r  c                       e Zd ZdZdedef fdZ	 ddej                  dej                  dej                  dz  d	e	e
eej                  ej                  f   f   d
edz  dee   deej                  ej                  dz  f   fdZ xZS )Gemma4TextAttentionz=Multi-headed attention from 'Attention Is All You Need' paperrD   rM   c                    t         |           t        |d      r|j                  |   nd | _        || _        || _        | j                  dk(  | _        | j                  r|j                  nd | _        | j                  s|j                  r|j                  n|j                  | _
        |j                  xr | j                   | _        | j                  r|j                  n|j                  }|j                  |z  | _        d| _        | j
                  j$                  | _        |j&                  dk7  | _        | j
                  j*                  t-        | j
                  dd      z
  }||cxk\  xr dk\  nc | _        |j                  d | }| j.                   xr6 |t1        |      dz
  |d d d   j3                  |j                  |         z
  k(  | _        t7        j8                  |j:                  |j                  | j                  z  |j<                  	      | _        tA        | j                  |jB                  
      | _"        | j.                  stA        | j                  |jB                  
      | _#        tA        | j                  |jB                  d      | _$        t7        j8                  |j:                  || j                  z  |j<                  	      | _%        | j                  s9t7        j8                  |j:                  || j                  z  |j<                  	      nd | _&        t7        j8                  |j                  | j                  z  |j:                  |j<                  	      | _'        y )Nr  rR   r   r  r  r   r?   r   rz   )r   r  Fr@  )(r   r   hasattrr  r  rD   rM   r  rT   r  r   attention_k_eq_vuse_alternative_attentionnum_global_key_value_headsr  r   num_key_value_groupsr  r  use_bidirectional_attentionr  r  r  r  lenindexstore_full_length_kvr   r   r   attention_biasr   r   r  r  r  r  r   r   r  )r   rD   rM   r  r  prev_layersr   s         r]   r   zGemma4TextAttention.__init__  s   ;B6=;Y&,,Y7_c"//-@@7;f33D6:oo&J`J`..flfufu)/)@)@)XEX&151O1OF--U[UoUo 	 %+$>$>BU$U!!%!>!>;;uD %)KK$A$AGDKKYoqrDs$s!"+/H"MA"M(()C*CD(,(?(?$? %/IQTU`QadeQehsbDi

%""9-
.R/ E/! ii : :T]] JQWQfQf
 $6;N;NO &&'DMMv?R?RSDK'6;N;N[`aDK))""$7$--$GfNcNcDK
 55 		&,,.ADMM.QX^XmXmn K ii&&68J8JQWQfQf
r_   Nr   r   rF   rb   rG   rT  rJ   c                    |j                   d d }g |d| j                  }|\  }	}
| j                  |      j                  |      }| j	                  |      }t        ||	|
d      }|j                  dd      }| j                  rI|| j                     \  }}|j                  |j                        }|j                  |j                        }n| j                  |      j                  |      }| j                   | j                  |      j                  |      n|}| j                  |      }t        ||	|
d      }|j                  dd      }| j                  |      }|j                  dd      }|,| j                  s |j                  ||| j                         \  }}| j"                  r||f|| j                  <   t%        j&                  | j(                  j*                  t,              } || ||||f| j.                  r| j0                  nd| j2                  | j4                  d|\  }} |j6                  g |d j9                         }| j;                  |      }||fS )Nr   r*   )r  r?   rd  )r  r  rT   )r   r   r   r   r  r:   rI  r  r  r   r   r   r   r  r  updaterM   r  r   r  rD   r  r;   r  r  r  rT   r   r   r  )r   r   r   rF   rb   rG   rT  r  r   r   r   r   r   r   r  r  r  s                    r]   r   zGemma4TextAttention.forward?  sR    $))#2.88b8$--8&S{{=166|D{{<0+L#sRST#--a3
 ""'7'H$J#|':':;J'??<+>+>?L]388FJLPKKLc4;;}5::<HisLZ0J-j#sRSTJ#--a3J;;|4L'11!Q7L&t/F/F'6'='=j,X\XfXf'g$J$$0:L0HT__-(?(M(MKK,,.E)
 %8
%
 /3mmD**LL..
%
 
%
!\ *k));;;;FFHkk+.L((r_   r   )rd   re   rf   rg   rB   r   r   rk   rl   rh   ri   rj   r   r   r   r   r   r   s   @r]   r  r    s    G/
/ /
C /
n )-=)||=) #\\=) t+	=)
 sE%,,*D$EEF=) =) -.=) 
u||U\\D00	1=)r_   r  c                   $     e Zd Zdef fdZ xZS )Gemma4TextExpertsrD   c                     t         |           |j                  | _        |j                  | _        t
        |j                     | _        y r   )r   r   num_expertsmoe_intermediate_sizeintermediate_dimr   hidden_activationr+  r/  s     r]   r   zGemma4TextExperts.__init__  s<    !-- & < <V556r_   )rd   re   rf   rB   r   r   r   s   @r]   r  r    s    7/ 7 7r_   r  c                   z     e Zd Zdef fdZdej                  deej                  ej                  f   fdZ xZ	S )Gemma4TextRouterrD   c                 (   t         |           || _        |j                  | _        | j                  dz  | _        |j
                  | _        t        | j                  | j                  d      | _        t        j                  |j                  |j                  d      | _        t        j                  t        j                  | j                              | _        t        j                  t        j                  |j                              | _        y )Nr   Fr@  rz   )r   r   rD   r   scalar_root_sizer  r  r   r  r   r   r!  projr   rk   r\  scaleper_expert_scaler/  s     r]   r   zGemma4TextRouter.__init__  s    !-- $ 0 0$ 6&&!$"2"2US	IIf00&2D2D5Q	\\%**T-=-=">?
 "UZZ8J8J-K Lr_   r   rJ   c                    | j                  |      }|| j                  z  | j                  z  }| j                  |      }t        j
                  j                  |dt        j                        }t        j                  || j                  j                  d      \  }}||j                  dd      z  }|| j                  |   z  }|||fS )Nr   r   )r  r   Trw  )r  r*  r(  r)  r   r   r   rk   r   topkrD   top_k_expertssumr+  )r   r   expert_scoresrouter_probabilitiestop_k_weightstop_k_indexs         r]   r   zGemma4TextRouter.forward  s    		-0%

2T5J5JJ		-0!}}44]RWR_R_4` &+ZZ kk''&
"{ 	**r4*@@ &(=(=k(JJ#]K??r_   )
rd   re   rf   rB   r   rk   rl   rj   r   r   r   s   @r]   r&  r&    s>    
M/ 
M@U\\ @eELL%,,<V6W @r_   r&  c                   0    e Zd Zdeez  def fdZ	 	 	 	 	 	 ddej                  dej                  de	e
eej                  ej                  f   f   dz  dej                  d	ej                  dz  d
ej                  dz  dedz  dej                  fdZ xZS )Gemma4TextDecoderLayerrD   rM   c                    t         |   ||       t        ||      | _        t	        ||      | _        | j                  dt        j                  d             |j                  | _	        | j                  rt        |j                     | _        t        j                  | j                  | j                  d      | _        t        j                  | j                  | j                  d      | _        t%        | j                  |j&                        | _        |j*                  | _        | j*                  rt-        |      | _        t1        |      | _        t%        | j                  |j&                        | _        t%        | j                  |j&                        | _        t%        | j                  |j&                        | _        y y )Nr  layer_scalarr?   Frz   r  )r   r   r  rO  r  r  r   rk   r\  hidden_size_per_layer_inputr   r$  r+  r   r   r   per_layer_input_gateper_layer_projectionr   r  post_per_layer_input_normenable_moe_blockr&  routerr  expertspost_feedforward_layernorm_1post_feedforward_layernorm_2pre_feedforward_layernorm_2r   s      r]   r   zGemma4TextDecoderLayer.__init__  sX   +,FiP 3^UZZ];+1+M+M(++ !9!9:DK(*		$2B2BDDdDdkp(qD%(*		$2R2RTXTdTdkp(qD%-:4;K;KQWQdQd-eD* & 7 7  *62DK,V4DL0=d>N>NTZTgTg0hD-0=d>N>NTZTgTg0hD-/<T=M=MSYSfSf/gD, !r_   Nr   per_layer_inputrb   r   rF   rH   rG   rJ   c           
      &   |}	| j                  |      } | j                  d||||||d|\  }}
| j                  |      }|	|z   }|}	| j                  |      }| j	                  |      }| j
                  r| j                  |      }|	j                  d|	j                  d         }| j                  |      \  }
}}| j                  |      }| j                  |||      }|j                  |	j                        }| j                  |      }||z   }| j                  |      }|	|z   }| j                  rP|}	| j                  |      }| j!                  |      }||z  }| j#                  |      }| j%                  |      }|	|z   }|| j&                  z  }|S )N)r   r   rF   rb   rH   rG   r   rS   )r  rO  r  r  r  r=  r@  r   r   r>  rB  r?  rA  r  r9  r:  r+  r;  r<  r7  )r   r   rC  rb   r   rF   rH   rG   rT  r3  rX   hidden_states_1hidden_states_flatr2  r3  hidden_states_2s                   r]   r   zGemma4TextDecoderLayer.forward  s    !,,];)4>> 
' 3)-%+
 
q 55mD =0 66}E/  "??NO "*!1!1"hnnR6H!I,0KK8J,K)A}k">>?QRO"ll?KWO-55hnnEO"??PO ,o=M77F =0++$H 55mDM KK6M)O;M 55mDM ::=IM$}4M***r_   )NNNNNN)rd   re   rf   rB   rC   r   r   rk   rl   rh   ri   rj   r  r   r   r   r   s   @r]   r5  r5    s    h/2DD hQT h0 )-PT,0.204(,9||9 9 sE%,,*D$EEFM	9
 #\\9 t+9 &&-9 9 
9r_   r5  c                       e Zd Zy)Gemma4TextScaledWordEmbeddingNr   rS   r_   r]   rI  rI    r   r_   rI  c                   J    e Zd Zg dZdZdZ ej                         d        Zy)Gemma4PreTrainedModel)r5  r  rW  rK  )imagetextvideoaudioNc                 	   t        j                  |       t        |t              r t	        j
                  |j                         y t        |t              rd}d}|j                  dz  }t        j                  ||z        t        |dz
  d      z  }|t        j                  t        j                  |      | z        z  }t	        j                  |j                   |j#                  d      j#                  d             y t        |t$              rJt	        j&                  |j(                  |j*                         t	        j,                  |j.                         y t        |t0              r|j2                  j5                         D ]  \  }}d|i}	|dk(  r|j6                  |   dk(  rd	|	d
<    ||j8                  fi |	\  }
}t	        j                  t;        || d      |
       t	        j                  t;        || d      |
        y t        |t<              r|j6                  dk7  rt>        |j6                     n|j@                  } ||j8                        \  }}t	        j                  |jB                  |       t	        j                  |jD                  |       y t        |tF              r+t	        j&                  |jH                  |jJ                         y t        |tL              r?t	        j
                  |jN                         t	        j
                  |jP                         y t        |tR              r[| j8                  jT                  }t	        jV                  |jX                  d|       t	        jV                  |jZ                  d|       y t        |t\              r t	        j
                  |j^                         y t        |t`              r|jb                  rt	        j&                  |jd                  tg        d              t	        j&                  |jh                  tg        d             t	        j&                  |jj                  tg        d              t	        j&                  |jl                  tg        d             y t        |tn              rV|j8                  jp                  r?t	        j,                  |jr                         t	        j
                  |jt                         y y y )Nr   r   r*   r?   r   r  rQ   r  r  r  r  r  r  rd  )meanstdr}   );r   _init_weightsr  rW  initones_r]  r   r   r   r   r   rk   r   r   copy_r   r   r   	constant_r   r   zeros_r   r  r  itemsr  rD   r  r  r   r  r  original_inv_freqrI  embed_scalescalar_embed_scaler&  r*  r+  r  initializer_rangenormal_gate_up_projr  r5  r7  rv   r   r|   r   r~   r   r   Gemma4VisionModelstandardizestd_bias	std_scale)r   moduler   r   r   r   r   r  r	  r
  r  rX   rope_fnbuffer_valuerR  s                  r]   rS  z#Gemma4PreTrainedModel._init_weights  s   %%f-f78JJv667 @AM#M#//14N&*hh}}/L&MPSTbefTfhiPj&j#*UYYu||N7SWnVn7n-ooNJJv,,n.F.Fq.I.S.STU.VW 45NN6>>6+K+KLKK,,- 9:,2,@,@,F,F,H ^(
L'3Z&@#!11f6F6Fz6RVd6d:K'7#/#UAT#U q

76j\+CDmT

76j\9K+LM}]^  ;< ##y0 $F$4$45;; 
 &fmm4OL!JJv5JJv//> =>NN6--v/H/HI 01JJv||$JJv../ 12++//CLL,,3C@LL))= 67JJv**+ 566;U;UNN6++eEl];NN6++U5\:NN6,,uU|m<NN6,,eEl; 12v}}7P7PKK(JJv''( 8Q2r_   )	rd   re   rf   _no_split_modulesinput_modalities_can_record_outputsrk   r   rS  rS   r_   r]   rK  rK    s2     ;U]]_2) 2)r_   rK  zAThe base Gemma 4 language model without a language modeling head.custom_introc                       e Zd ZU eed<    eed      eedZ	def fdZ
dej                  dz  dej                  dz  d	ej                  fd
Z	 ddej                  dej                  dz  d	ej                  fdZeee	 	 	 	 	 	 	 ddej$                  dz  dej                  dz  dej$                  dz  dedz  dej(                  dz  dej                  dz  dedz  dee   d	efd                     Z xZS )Gemma4TextModelrD   r   )r  )router_logitsr   
attentionsc           
         t         |   |       t        j                  t	        |j
                        D cg c]  }t        ||       c}      | _        t        |      | _	        t        | j                  j                        | _        |j                  | _        | j                  rt        |j                   |j
                  |j                  z  | j"                  |j                  dz        | _        d| _        t        j(                  |j*                  |j
                  |j                  z  d      | _        |j*                  dz  | _        t1        |j                  |j2                        | _        g | _        t9        | j                        D ]K  \  }}|j:                  j<                  s| j6                  j?                  dD cg c]
  }d	| d
|  c}       M y c c}w c c}w )Nrm  )r[  g;f?Frz   r   r8  )r   r   r  r  zlayers.z.self_attn.) r   r   r   r  r  r  r5  r  r  r  r  rD   r  unique_layer_typesr9  rI  vocab_size_per_layer_inputpadding_idxembed_tokens_per_layerper_layer_input_scaler   r   per_layer_model_projection per_layer_model_projection_scaler   r  per_layer_projection_norm"_keys_to_ignore_on_load_unexpected	enumeraterO  r  extend)r   rD   rM   r  layernamer   s         r]   r   zGemma4TextModel.__init__N  s    mmHMfNfNfHgh9#FI6h
 4F;"%dkk&=&=">
 ,2+M+M(++*G11((6+M+MM  ">>C	+D' *3D&.0ii""((6+M+MM/D+
 5;4F4F4LD1-:6;];]cicvcv-wD* 35/!$++. 	HAu1177>>@hiwqcTF3i	7 i< js   GG"
	input_idsNrE   rJ   c                    | j                   st        d| j                         |t        j                         5  |dddddddf   | j
                  j                  ddddddf   | j                  j                  dz  z  k(  j                  d      j                         dddf   }	 |j                  |j                  dd       }	 ddd        | j                  |      j                  g |j                  | j                  j                  | j                    S # t        $ r t        d      w xY w# 1 sw Y   nxY w)a  Compute the token-identity component of Per-Layer Embeddings (PLE).

        Looks up `input_ids` in `embed_tokens_per_layer` (a scaled embedding that multiplies
        by `sqrt(hidden_size_per_layer_input)`) and reshapes the packed output from
        `[batch, seq, num_hidden_layers * hidden_size_per_layer_input]` to
        `[batch, seq, num_hidden_layers, hidden_size_per_layer_input]`.

        If only `inputs_embeds` is provided (no `input_ids`), reverses the main embedding
        to recover `input_ids` for the PLE lookup.
        z}Attempting to call get_per_layer_inputs() from a model initialized with a config that does not support per-layer embeddings. Nrm  r	   r   r*   a)  It seems like you tried to call `forward` from `inputs_embeds` without providing `input_ids`, and that the `inputs_embeds` you provided do not exactly match the embedding weights. Since Gemma4 needs to reverse the embedding to compute another embedding, make sure you provide exact `inputs_embeds`)r9  RuntimeErrorrD   rk   r   embed_tokensr  r   r  nonzeror   r   rt  r   r  )r   r~  rE   s      r]   get_per_layer_inputsz$Gemma4TextModel.get_per_layer_inputsr  sU    //**.++8    &aD!m4,,33D$14DEH_H_adHdde SQSZWYq!t%  )}/B/B2A/F GI$ >t**95== 
__
KK))
 ,,
 	
 $ &r  s   A1D9-D!!D66D99Eper_layer_inputsc                 V   | j                   st        d| j                         | j                  |      | j                  z  } |j
                  g |j                  dd | j                  j                  | j                    }| j                  |      }||S ||z   | j                  z  S )a  Compute the context-aware component of PLE and combine with token-identity.

        Projects `inputs_embeds` through `per_layer_model_projection` (Linear), scales by
        `1/sqrt(hidden_size)`, reshapes to `[batch, seq, num_layers, ple_dim]`, and normalizes
        with `per_layer_projection_norm` (RMSNorm).

        If `per_layer_inputs` (the token-identity component from `get_per_layer_inputs()`)
        is provided, combines both: `(context_projection + token_identity) * (1/sqrt(2))`.
        If `per_layer_inputs` is None (e.g. for multimodal inputs where input_ids are not
        available), returns just the context projection.
        zAttempting to call project_per_layer_inputs() from a model initialized with a config that does not support per-layer embeddings. Nr   )
r9  r  rD   rv  rw  r   r   r  rx  ru  )r   rE   r  r;  s       r]   project_per_layer_inputsz(Gemma4TextModel.project_per_layer_inputs  s      //226++@ 
  $>>}MPTPuPuu;3;;  
  "% 
KK)) 
 ,, 

  $==>RS#''$'774;U;UUUr_   rF   rH   rG   	use_cacherT  c           
      <   |du |duz  rt        d      ||t        d      || j                  |      }| j                  r&|| j                  ||      }| j	                  ||      }|r|t        | j                        }|V||j                         nd}	t        j                  |j                  d   |j                        |	z   }|j                  d      }t        |x}
t              s)| j                  ||||d}t        di |t!        di |d	}
|}i }| j"                  D ]  }| j%                  |||      ||<    |j'                  d
t)                     }t+        | j,                  d| j                  j.                         D ]\  \  }}||dddd|ddf   nd} |||f||| j                  j0                  |      |
| j                  j0                  |      ||d|}^ | j3                  |      }t5        |||j7                  dd      r|      S d      S )ak  
        per_layer_inputs (`torch.Tensor`, *optional*):
            Pre-computed per-layer input text embeddings of shape `(batch_size, sequence_length, num_hidden_layers,
            hidden_size_per_layer_input)`. When provided, these are used directly instead of being computed from `input_ids`
            via `get_per_layer_inputs()` in the text model. If calling the `forward` with `inputs_embeds` instead of `input_ids`,
            you should probably precompute them and forward them along `inputs_embeds`, otherwise recomputing them needs
            to reverse the main embedding, which is expensive.
        N:You must specify exactly one of input_ids or inputs_embeds<You cannot specify per_layer_inputs if input_ids is provided)rD   r   r?   r   rL   rP   rb   )rb   r   rF   rH   rG   return_shared_kv_statesF)r  rG   rb   rS   )r{  r  r9  r  r  r   rD   get_seq_lengthrk   r   r   r   r   r  rh   r   r   rq  r  popr   rz  r  r  r  r  rq   get)r   r~  rF   rH   rG   rE   r  r  rT  past_seen_tokenscausal_mask_mappingrU   r   r   r  rb   r  r  rC  s                      r]   r   zGemma4TextModel.forward  s}   , -t";<YZZ %5%A[\\  --i8M++'#'#<#<Y#V #<<]L\]0*$++>OCRC^==?de <<(;(;A(>}G[G[\_ooL'11!4L ?-F ++!."0#2 ,K #5"C{"C%F%U%U# & 11 	gJ.2oom\[e.f
+	g "::&8(*E !*$++6U8U8U*V W 	A}>N>Z.q!Qz:`dO)	 "2$78O8OPQ8R$S24;;3J3J13MN) /	 	M	 		-0,++17<UW\1]-
 	
 dh
 	
r_   r   )NNNNNNN)rd   re   rf   rB   rm   r(   r&  r5  r  ri  r   rk   rl   r  r  r&   r)   r    r  r   r  boolr   r   rq   r   r   r   s   @r]   rm  rm  E  s   '(8B/)"/ "H*
ellT.A *
RWR^R^aeRe *
jojvjv *
^ 15!V||!V  ,,-!V 
	!VF   .2.204(,2604!%Y
##d*Y
 t+Y
 &&-	Y

 Y
 ((4/Y
  ,,-Y
 $;Y
 +,Y
 
'Y
    Y
r_   rm  z>The base Gemma 4 language model with a language modeling head.c                   4   e Zd ZdZee	 	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  de	dz  dej                  dz  dej                  dz  d	edz  d
eej                  z  dej                  dz  dee   defd              Zy)Gemma4ForCausalLMmodelNr~  rF   rH   rG   rE   labelsr  logits_to_keepr  rT  rJ   c
                 4    | j                   d||||||	|d|
}|j                  }t        |t              rt	        | d      n|}| j                  |dd|ddf         }| j                  j                  G|| j                  j                  z  }t        j                  |      }|| j                  j                  z  }d}| | j                  ||| j                  fi |
}t        |||j                  |j                  |j                  |j                         S )a  
        per_layer_inputs (`torch.Tensor`, *optional*):
            Pre-computed per-layer input text embeddings of shape `(batch_size, sequence_length, num_hidden_layers,
            hidden_size_per_layer_input)`. When provided, these are used directly instead of being computed from `input_ids`
            via `get_per_layer_inputs()` in the text model. If calling the `forward` with `inputs_embeds` instead of `input_ids`,
            you should probably precompute them and forward them along `inputs_embeds`, otherwise recomputing them needs
            to reverse the main embedding, which is expensive.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, Gemma4ForCausalLM

        >>> model = Gemma4ForCausalLM.from_pretrained("google/gemma-4-E2B-it")
        >>> tokenizer = AutoTokenizer.from_pretrained("google/gemma-4-E2B-it")

        >>> prompt = "What is your favorite condiment?"
        >>> inputs = tokenizer(prompt, return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
        >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "What is your favorite condiment?"
        ```)r~  rF   rH   rG   rE   r  r  N)losslogitsrG   r   ro  rb   rS   )r  r  r  r   slicelm_headrD   final_logit_softcappingrk   r   loss_function
vocab_sizero   rG   r   ro  rb   )r   r~  rF   rH   rG   rE   r  r  r  r  rT  outputsr   slice_indicesr  r  s                   r]   r   zGemma4ForCausalLM.forward$  s"   P 2< 	2
)%+'-	2
 	2
  118B>SV8W~ot4]kmA}a,?@A;;..:dkkAAAFZZ'FdkkAAAF%4%%ffdooPPD+#33!//))$55
 	
r_   )	NNNNNNNr   N)rd   re   rf   base_model_prefixr!   r    rk   r  rl   r   r  r  r   r   r   ro   r   rS   r_   r]   r  r     s    .2.204(,26*.!%-.04E
##d*E
 t+E
 &&-	E

 E
 ((4/E
   4'E
 $;E
 ell*E
  ,,-E
 +,E
 
&E
  E
r_   r  c                   ,    e Zd ZU dZeed<   dZdZee	dZ
def fdZdej                  dej                  fd	Zee ed
      	 ddej                  dej                  dz  dee   deej                  ej*                  f   fd                     Z xZS )Gemma4AudioModelznAn audio encoder based on the [Universal Speech Model](https://huggingface.co/papers/2303.01037) architecture.rD   r   zmodel.audio_towerr   ro  c           	         t         |   |       || _        t        |      | _        t        |      | _        t        j                  t        |j                        D cg c]  }t        ||       c}      | _        t        j                  |j                  |j                  d      | _        | j#                          y c c}w )NTrz   )r   r   rD   r  subsample_conv_projectionr   rel_pos_encr   r  r  r  rK  r  r   r   output_proj_dimsoutput_proj	post_initr   s      r]   r   zGemma4AudioModel.__init__y  s     )KF)S&;FCmmBGH`H`BabYfi0b
 99V%7%79P9PW[\	 cs   B?mask_4drJ   c                    |j                   \  }}}}|j                  }| j                  j                  }| j                  j                  dz
  }| j                  j
                  }||z   dz
  |z  }	|	|z  }
|
|z
  }t        j                  |d|d|fd      }|j                  |d|	||
      }t        j                  |||fd      }t        j                  |	|      |z  }t        j                  ||z   |z   |      }|dddf   |dddf   z   }|dddddddf   j                  |dd|d      }|j                  d|      S )z
        Convert a standard 4D attention mask `[batch_size, 1, seq_len, seq_len]` to the 5D blocked format
        `[batch_size, 1, num_blocks, chunk_size, context_size]` expected by the chunked local attention,
        r?   r   F)valuer   Nr   )r   r   rD   r   r   r   r   r   r   rk   r   r  gather)r   r  r   rX   r   r   r   r   r   r   padded_seq_len
pad_amountmask_5dblock_startsoffsets
kv_indicess                   r]   _convert_4d_mask_to_blocked_5dz/Gemma4AudioModel._convert_4d_mask_to_blocked_5d  sN   
 %,MM!
Aw[[55
;;==A![[@@
*Q.:=
#j0#g-
%%!ZJ!?uM//*aZX%%"24F!GuU||Jv>K,,z,<<?QQZ`a!!T'*WT1W-==
dAtQ 67>>z1bR\^`a
~~b*--r_   z&Encodes audio features to soft tokens.rj  NrF   rT  c           	         | j                  ||      \  }}| j                  |      }t        | j                  ||t	        | j                  j
                  dz
  | j                  j                  f            }|| j                  |      }| j                  d | j                  j                   D ]  } ||f||d|} | j                  |      }t        ||      S )Nr?   )rD   rE   rF   rO   )rF   r   )r  rF   )r  r  r   rD   r>   r   r   r  r  r  r  rs   )r   r   rF   rT  r   output_maskr   encoder_layers           r]   r   zGemma4AudioModel.forward  s     &*%C%CNTb%c"{"..}=2;;'&:33a79\9\]	
 %!@@PN![[)H4;;+H+HI 	M)-$7 	M	 ((7%Vabbr_   r   )rd   re   rf   rg   r@   rm   main_input_namer  rK  r   ri  r   rk   rl   r  r&   r)   r    r   r   rj   rt   r   r   r   s   @r]   r  r  n  s    x&O+)*
0 .ell .u|| .6  !IJ /3cc t+c +,	c
 
u||U---	.c K   cr_   r  c                        e Zd ZdZeZeedZdef fdZ	e
e ed      dej                  dej                  d	ee   d
efd                     Z xZS )r`  zThe Gemma 4 Vision Encoder.r  rD   c                    t         |   |       t        |      | _        t	        |      | _        t        |      | _        | j                  j                  rr| j                  dt        j                  | j                  j                               | j                  dt        j                  | j                  j                               | j                          y )Nrb  rc  )r   r   rW  patch_embedderr  encoderrq  poolerrD   ra  r   rk   emptyr   r  r/  s     r]   r   zGemma4VisionModel.__init__  s     7?*62(0;;""  U[[9P9P-QR  ekk$++:Q:Q.RSr_   z1Encodes image pixels to soft tokens from patches.rj  rk  r^  rT  rJ   c                    | j                   j                  }|j                  d   ||z  z  }|dk(  j                  d      }| j	                  |||      } | j
                  d|| |d|}| j                  |j                  |||      \  }	}
|	|
   }	| j                   j                  r8|	| j                  j                         z
  | j                  j                         z  }	|	j                  |j                        }	t        |	      S )a  
        pixel_values (`torch.FloatTensor` or `list[torch.FloatTensor]`):
            The images to encode. Either a single `[batch, channels, height, width]` tensor
            (all images same size) or a list of `[1, channels, height, width]` tensors (different sizes).
        pixel_position_ids (`torch.LongTensor` of shape `(batch_size, max_patches, 2)`):
            The patch positions as (x, y) coordinates in the image. Padding patches are indicated by (-1, -1).
        r   r   )rE   rF   r^  )r   r^  r_  r  r  rS   )rD   pooling_kernel_sizer   r  r  r  r  r  ra  rb  r   rc  r   r   r   )r   rk  r^  rT  r  r  r_  rE   r  r   pooler_masks              r]   r   zGemma4VisionModel.forward  s     #kk==$**2.3FI\3\]/25::r:B++L:LN_` 
'--1
 	
 &*[[ 221/'	 &1 &
"{ &k2 ;;""*T]]-@-@-BBdnnFZFZF\\M%(()<)<=&GGr_   )rd   re   rf   rg   rC   rD   r  r  ri  r   r&   r)   r    rk   r  r  r   r   r   r   r   r   s   @r]   r`  r`    s    %F1+

1 
  !TU)H'')H ",,)H +,	)H
 
!)H V   )Hr_   r`  c                   f     e Zd Zdeez  def fdZdej                  dej                  fdZ	 xZ
S )Gemma4MultimodalEmbeddermultimodal_configtext_configc                     t         |   ||       | `| `| `| `| `| `t        |d|j                        | _
        t        | j                  | j                  d      | _        y )Nr  Fr@  )r   r   re  hard_embedding_normsoft_embedding_normvocab_offsetr  embedding_post_projection_normr  r   multimodal_hidden_sizer   r  embedding_pre_projection_norm)r   r  r  r   s      r]   r   z!Gemma4MultimodalEmbedder.__init__	  sp     	*K8N$$O/&-.?ASUfUrUr&s#-:4;V;V\`\d\dqv-w*r_   rE   rJ   c                 F    | j                  |      }| j                  |      S )a:  Embeds token ids or soft tokens for multimodal content into language model space.
        Args:
            inputs_embeds: A torch.Tensor containing the soft tokens to embed.
        Returns:
            A torch.Tensor of embeddings with shape `[batch_size, seq_len, self.config.text_config.hidden_size]`.
        )r  embedding_projection)r   rE   embs_normeds      r]   r   z Gemma4MultimodalEmbedder.forward  s%     88G((55r_   )rd   re   rf   r@   rC   rB   r   rk   rl   r   r   r   s   @r]   r  r    s>    x,/AAx &x$6U\\ 6ell 6r_   r  token_type_idsimage_group_idsc           
      V    | ydt         dt         dt         dt         dt        f
fd}|S )z
    This function adds the correct offsets to the `q_idx` and `kv_idx` as the torch API can only accept lengths,
    not start and end indices.
    N	batch_idxhead_idxq_idxkv_idxrJ   c                    	j                   d   }|j                  |dz
        }|j                  |dz
        }	| |f   }	| |f   }t        j                  ||k  |d      }t        j                  ||k  |d      }||k(  |dk\  z  S )Nr   r?   )r   r   )r   r   rk   rf  )
r  r  r  r  r   q_idx_clampedkv_idx_clampedq_groupkv_groupr  s
            r]   
inner_maskz0token_type_ids_mask_function.<locals>.inner_mask3  s    $**2.
 
Q7*q.9 ")]":;"9n#<=++ej0'2>;;v
2HbA8#155r_   )r   r  )r  r  r  s    ` r]   token_type_ids_mask_functionr  '  s>     6c 6S 6 6c 6d 6 r_   mm_token_type_idsr   c                    | j                  |      } | dk(  | dk(  z  }t        j                  |dd      }d|d<   || z  }t        j                  |j	                         d      dz
  }t        j
                  ||d      }|S )Nr?   r*   r   )shiftsdimsFrb  r   )r   rk   rollcumsumr   rf  )r  r   	is_visionis_prev_visionnew_vision_startsvision_group_idsrI   s          r]   get_block_sequence_ids_for_maskr  D  s    ),,V4"a',=,BCIZZ	!"=N"N6!^O3||$5$9$9$;CaGY0@"Er_   z
    The base Gemma 4 model comprising a vision backbone, an audio backbone, and a language model without a
    language modeling head.
    c            $           e Zd Zdef fdZd Zd Ze ed      	 dde	j                  d	e	j                  dz  d
ee   defd              Ze ed      	 dde	j                  de	j                  dz  d
ee   defd              Z	 	 d de	j                  dz  de	j                  dz  dee	j$                  e	j$                  e	j$                  f   fdZeee	 	 	 	 	 	 	 	 	 	 	 	 	 	 d!de	j                  dz  de	j                  dz  de	j                  dz  de	j                  dz  de	j*                  dz  de	j*                  dz  de	j                  dz  dedz  de	j                  dz  de	j                  dz  dedz  d	e	j                  dz  de	j                  dz  de	j*                  dz  d
ee   def d                     Ze ed      de	j*                  de	j*                  d
ee   deez  fd              Z xZS )"Gemma4ModelrD   c                    t         |   |       |j                  t        j                  |j                        nd | _        |j                   t        |j                  |j                        nd | _        |j                  t        j                  |j                        nd | _
        |j                  &t        |j                  |j                        | _        y d | _        y r   )r   r   vision_configr+   from_configvision_towerr  r  embed_visionaudio_configaudio_towerembed_audior/  s     r]   r   zGemma4Model.__init__W  s     KQK_K_KkI11&2F2FGqu ##/ %V%9%96;M;MN 	
 JPI\I\Ih9001D1DEnr "". %V%8%8&:L:LM 	  	r_   c                 .    | j                   j                  S r   language_modelrt  r   s    r]   get_per_layer_input_embeddingsz*Gemma4Model.get_per_layer_input_embeddingsf  s    ""999r_   c                 &    || j                   _        y r   r  r   r  s     r]   set_per_layer_input_embeddingsz*Gemma4Model.set_per_layer_input_embeddingsi  s    5:2r_   zOProjects the last hidden state from the vision model into language model space.rj  Nrk  image_position_idsrT  rJ   c                 v     | j                   d||d|}|j                  }| j                  |      |_        |S )z
        image_position_ids (`torch.LongTensor` of shape `(batch_size, max_patches, 2)`, *optional*):
            The patch positions as (x, y) coordinates in the image. Padding patches are indicated by (-1, -1).
        rk  r^  rE   rS   )r  r  r  pooler_output)r   rk  r  rT  vision_outputsr  s         r]   get_image_featureszGemma4Model.get_image_featuresl  sV     +** 
%1
 

 +<<'+'8'8GX'8'Y$r_   zQProjects the last hidden state from the vision encoder into language model space.pixel_values_videosvideo_position_idsc                     |j                  dd      }|j                  dd      } | j                  d||d|}|j                  }| j                  |      |_        |S )a9  
        video_position_ids (`torch.LongTensor` of shape `(num_videos, num_frames, max_patches, 2)`, *optional*):
            2D patch position coordinates from the video processor, with `(-1, -1)` indicating padding.
            Passed through to the vision encoder for positional embedding computation.
        r   r?   r  r   rS   )flattenr  r  r  r  )r   r  r  rT  r  r  s         r]   get_video_featureszGemma4Model.get_video_features  s|     299!Q?/771=*** 
,1
 

 +<<'+'8'8GX'8'Y$r_   r~  rE   c                 &   |M|| j                   j                  k(  }|| j                   j                  k(  }|| j                   j                  k(  }n>| | j	                         t        j                  | j                   j                  t
        j                  |j                              k(  j                  d      }| | j	                         t        j                  | j                   j                  t
        j                  |j                              k(  j                  d      }| | j	                         t        j                  | j                   j                  t
        j                  |j                              k(  j                  d      }|||fS )a  
        Obtains mask for multimodal placeholders (replaced by soft tokens) and hard text tokens.

        Masks will be obtained from `mm_token_type_ids`, `input_ids`, or `inputs_embeds` as available and in that
        precedence order. If passing `input_ids` or `inputs_embeds`, the image mask will be derived using
        `config.image_token_id`. Same goes for audio and video masks

        Args:
            input_ids: A tensor containing the hard token IDs from the text tokenizer.
            inputs_embeds: A tensor containing the embeddings for all hard text tokens.

        Returns:
            image_mask, video_mask, audio_mask
        )r   r   r   )
rD   image_token_idvideo_token_idaudio_token_idget_input_embeddingsrk   r   r~  r   r  )r   r~  rE   special_image_maskspecial_video_maskspecial_audio_masks         r]   get_placeholder_maskz Gemma4Model.get_placeholder_mask  sP   &  !*dkk.H.H!H!*dkk.H.H!H!*dkk.H.H!H .4,,.LL!;!;5::VcVjVjk c"g  .4,,.LL!;!;5::VcVjVjk c"g  .4,,.LL!;!;5::VcVjVjk c"g  "#57IIIr_   r   rF   r!  rH   rG   r  r  r  c                 &   |du |
duz  rt        d      ||t        d      | j                  ||
      \  }}}||z  |z  }d}|
[|j                         }t        j                  || j
                  j                  j                  |      } | j                         |      }
|| j
                  j                         j                  r| j                  j                  j                  | j
                  j                  j                  ddf   }|j                  |
j                        }t        j                  |d   |j!                  ddd      |
      }| j                  j#                  ||      }|| j%                  ||d      j&                  }|j                  |
j                  |
j(                        }|j+                         }|j-                  d      j/                  |
      j                  |
j                        }t1        |
|   j3                         |j3                         k(  d	| d
|j4                  d           |
j7                  |j                  |
j                        |j                  |
j                              }
|| j9                  ||d      j&                  }|j                  |
j                  |
j(                        }|j+                         }|j-                  d      j/                  |
      j                  |
j                        }t1        |
|   j3                         |j3                         k(  d| d
|j4                  d           |
j7                  |j                  |
j                        |j                  |
j                              }
|+|(| j;                  ||d      }|j&                  }|j<                  }||j                  |j                           }|j+                         }|j-                  d      j/                  |
      j                  |
j                        }t1        |
|   j3                         |j3                         k(  d| d
|j4                  d   |j4                  d   z          |
j7                  |j                  |
j                        |j                  |
j                              }
|V||j?                         nd}t        j@                  |
j4                  d   |
j                        |z   }|j-                  d      }tC        |x} tD              s}| j
                  j                         |
|||d}!| j
                  j                         }"|"jF                  dk(  }#|#r'|	%tI        |	|
j                        }$tK        dd|$i|!} ntM        di |!}  | j                  d|| |||
|dd|}%tO        |%jP                  |%jR                  |%jT                  |%jV                  |nd|nd|%jX                        S )K  
        input_features_mask (`torch.FloatTensor]` of shape `(num_images, seq_length)`):
            The attention mask for the input audio.
        image_position_ids (`torch.LongTensor` of shape `(batch_size, max_patches, 2)`, *optional*):
            2D patch position coordinates from the image processor, with `(-1, -1)` indicating padding.
            Passed through to the vision encoder for positional embedding computation.
        video_position_ids (`torch.LongTensor` of shape `(num_videos, num_frames, max_patches, 2)`, *optional*):
            2D patch position coordinates from the video processor, with `(-1, -1)` indicating padding.
            Passed through to the vision encoder for positional embedding computation.
        per_layer_inputs (`torch.Tensor`, *optional*):
            Pre-computed per-layer input text embeddings of shape `(batch_size, sequence_length, num_hidden_layers,
            hidden_size_per_layer_input)`. When provided, these are used directly instead of being computed from `input_ids`
            via `get_per_layer_inputs()` in the text model. If calling the `forward` with `inputs_embeds` instead of `input_ids`,
            you should probably precompute them and forward them along `inputs_embeds`, otherwise recomputing them needs
            to reverse the main embedding, which is expensive.
        Nr  r  r   r?   r   T)return_dictz6Image features and image tokens do not match, tokens: z, features: r   z6Video features and video tokens do not match, tokens: z6Audio features and audio tokens do not match, tokens: r   rL   visionrI   )r  rF   rH   rG   rE   r  r  )r  rG   r   ro  image_hidden_statesaudio_hidden_statesrb   rS   )-r{  r  r  rk   rf  rD   r  pad_token_idr  get_text_configr9  r  r  r  r   r   r   r  r  r  r   r/  r   	expand_asr$   numelr   masked_scatterr  get_audio_featuresrF   r  r   r  rh   r  r  r^   r   ra   r  rG   r   ro  rb   )&r   r~  rk  r  r   rF   r!  rH   rG   r  rE   r  r  r  r  rT  
image_mask
video_mask
audio_maskmultimodal_maskllm_input_idspad_embeddingllm_inputs_embedsimage_featuresn_image_tokensvideo_featuresn_video_tokensaudio_outputaudio_featuresaudio_mask_from_encodern_audio_tokensr  r  rU   r  	use_bidirrI   r  s&                                         r]   r   zGemma4Model.forward  s   J -t";<YZZ %5%A[\\-1-F-FyR_-`*
J
$z1J>  %OO-M!KK9P9P9]9]_lmM7D557FM#(C(C(E(a(a //<<CCDKKD[D[DhDhjkDklM-001E1EFO %OI,FHZHZ[\^_acHdfs t#22GGWhi #!44\CUcg4hvvN+..}/C/C]EXEXYN (^^-N#--b1;;MJMMmNbNbcJ"j)//1^5I5I5KKHHX Y"((+,. *88m223^5F5F}G[G[5\M *!44#%7T 5 m  ,..}/C/C]EXEXYN (^^-N#--b1;;MJMMmNbNbcJ"j)//1^5I5I5KKHHX Y"((+,. *88m223^5F5F}G[G[5\M
 %*=*I22>CVdh2iL)77N&2&A&A#
 ,,C,F,F~G\G\,]^N'^^-N#--b1;;MJMMmNbNbcJ"j)//1^5I5I5KKHHX Y"((+n.B.B1.EEFH *88m223^5F5F}G[G[5\M
 CRC^==?de <<(;(;A(>}G[G[\_ooL'11!4L?-F++557!."0#2 ,K ++557K#??8KI.:%DEV_l_s_s%t"&C ''9'!'# '@&N+&N#%$%% 	
-.%+'	
 	
 )%77#33!//))2>2JPT2@2LRV$55
 	
r_   zPProjects the last hidden state from the audio encoder into language model space.c                     | j                   t        d       | j                   ||fddi|}| j                  |j                        |_        |S )a0  
        input_features (`torch.FloatTensor]` of shape `(num_images, seq_length, num_features)`):
            The tensors corresponding to the input audio.
        input_features_mask (`torch.FloatTensor]` of shape `(num_images, seq_length)`):
            The attention mask for the input audio.
        zAudio features were requested, but the model was initialized without an audio_config. Cannot process audio without an audio tower and audio embedder.r  Tr   )r  r{  r  r  r  )r   r   r!  rT  audio_outputss        r]   r  zGemma4Model.get_audio_featureso  sh     #R 
 )((9LiZ^ibhi&*&6&6]EdEd&6&e#r_   r   r  )NNNNNNNNNNNNNN)rd   re   rf   rA   r   r  r  r!   r    rk   r  r  r   r   r   r  r  rj   rt   r  r&   rl   r   r  ra   r   rs   r  r   r   s   @r]   r  r  P  s/   
| 
:; !rs 7;'' ",,t3 +,	
 
$ t & !tu 7;".. ",,t3 +,	
 
$ v 0 .226+J##d*+J ((4/+J 
u!1!153C3CC	D	+JZ   .2158<37.23704(,5926!%6:6:04d
##d*d
 ''$.d
 #..5	d

 ))D0d
 t+d
 #\\D0d
 &&-d
 d
 !++d2d
 ((4/d
 $;d
 ",,t3d
 ",,t3d
  ,,-d
  +,!d
" 
##d
    d
L !st #\\ +,	
 
'	' u r_   r  z
    The base Gemma 4 model comprising a vision backbone, an audio backbone, a language model, and a language modeling
    head.
    c            %       6    e Zd ZdZd Zd Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  dej                  dz  d	ej                  dz  d
ej                  dz  dej                  dz  dej                  dz  dej                  dz  de
dz  dej                  dz  dej                  dz  dej                  dz  dedz  deej                  z  dej                  dz  dee   def$dZe	 ddej                  dej                  dz  dee   fd       Ze	 	 ddedej                  d	ej                  dz  de
dz  dej                  dz  dej                  dz  dedz  defd       Z	 	 	 	 	 	 	 	 	 	 	 	 	 d  fd	Z xZS )!Gemma4ForConditionalGenerationr  c                 6    | j                   j                         S r   )r  r  r  s    r]   r  z=Gemma4ForConditionalGeneration.get_per_layer_input_embeddings  s    zz88::r_   c                 :    | j                   j                  |       y r   )r  r  r  s     r]   r  z=Gemma4ForConditionalGeneration.set_per_layer_input_embeddings  s    

11%8r_   Nr~  rk  r  r   rF   r!  rH   r  r  rG   r  rE   r  r  r  r  rT  rJ   c           
          | j                   di d|d|d|d|d|d|d|d|
d	|d
|d|d|d|d|d|	dd|}|j                  }t        |t              rt	        | d      n|}| j                  |dd|ddf         }| j                  j                         j                  x}||z  }t        j                  |      }||z  }d}|7 | j                  ||| j                  j                         j                  fi |}t        |||j                  |j                  |j                   |j"                  |j$                  |j&                        S )r  r~  rk  r  r   rF   r!  rH   rG   r  rE   r  r  r  r  r  r  TN)r  r  rG   r   ro  r  r  rb   rS   )r  r  r  r   r  r  rD   r  r  rk   r   r  r  ro   rG   r   ro  r  r  rb   )r   r~  rk  r  r   rF   r!  rH   r  r  rG   r  rE   r  r  r  r  rT  r  r   r  r  r  r  s                           r]   r   z&Gemma4ForConditionalGeneration.forward  s   H $** 

%
 !4
 *	

 *
 !4
 &
 ,
 0
 (
 .
 
  
  2
  2
  #
(  118B>SV8W~ot4]kmA}a,?@A'+{{'B'B'D'\'\\#i55FZZ'F55F%4%%ffdkk6Q6Q6S6^6^ibhiD+#33!//)) ' ; ; ' ; ;$55	
 		
r_   c                 >     | j                   j                  ||fi |S )a-  
        image_position_ids (`torch.LongTensor` of shape `(batch_size, max_patches, 2)`, *optional*):
            2D patch position coordinates from the image processor, with `(-1, -1)` indicating padding.
            Passed through to the vision encoder for positional embedding computation.
        )r  r  )r   rk  r  rT  s       r]   r  z1Gemma4ForConditionalGeneration.get_image_features  s$     -tzz,,\;MXQWXXr_   rD   is_first_iterationc                     | j                         ||||d}| j                         }	t        |	dd       dk(  }
|
r&|$t        ||j                        }t	        dd|i|S t        di |S )NrL   r  r  r   rI   rS   )r  r  r  r   r^   r   )rD   rE   rF   rG   rH   r  r6  rT  rU   r  r-  rI   s               r]   r   z8Gemma4ForConditionalGeneration.create_masks_for_generate  s     ,,.*,.(
 ,,.K)FMQYY	*6!@AR[h[o[o!p0 #5 
 )7;77r_   c                     t        |   |f|||||||
|d|}|s|s||d<   ||d<   ||d<   |	|d<   nd |d<   |s|j                  dd       }|S )N)rG   rE   rF   rH   r  r  r  r6  rk  r  r   r!  r  r  )r   prepare_inputs_for_generationr  )r   r~  rG   rE   rH   rk  r  r   rF   r!  r  r  r  r  r6  rT  model_inputsrX   r   s                     r]   r9  z<Gemma4ForConditionalGeneration.prepare_inputs_for_generation	  s    & w<
+')%))1
 
 Y+7L(2EL./-;L)*2EL./ 15L,- "  !3T:Ar_   )NNNNNNNNNNNNNNr   Nr   )NF)NNNNNNNNNTNNF)rd   re   rf   r  r  r  rk   r  r  rl   r   r  r   r   r   ro   r   r    r  r  r   rh   r   r9  r   r   s   @r]   r1  r1    s     ;9
 .2158<37.237046:6:(,5926*.!%-.04#N
##d*N
 ''$.N
 #..5	N

 ))D0N
 t+N
 #\\D0N
 &&-N
 ",,t3N
 ",,t3N
 N
 !++d2N
 ((4/N
   4'N
 $;N
  ell*!N
"  ,,-#N
$ +,%N
& 
&'N
`  7;Y''Y ",,t3Y +,	Y Y  26*/8 8||8 t+8 	8
 llT)8 !<<$.8 !4K8 
8 8B    . .r_   r1  )r  r  r1  r  rK  rm  r`  )r*   )r   collectionsr   collections.abcr   dataclassesr   	functoolsr   rk   r   torch.nnr   r    r
   rT  activationsr   cache_utilsr   r   configuration_utilsr   masking_utilsr   r   r   r   r   r   r   r   modeling_flash_attention_utilsr   modeling_outputsr   r   modeling_rope_utilsr   r   modeling_utilsr   r   processing_utilsr   utilsr   r    r!   r"   r#   r$   utils.genericr%   r&   r'   utils.output_capturingr(   r)   auto.modeling_autor+   gemma3.modeling_gemma3r,   r-   r.   r/   r0   r1   r2   gemma3n.modeling_gemma3nr3   r4   r5   r6   r7   r8   r9   r:   r;   llama.modeling_llamar<   mixtral.modeling_mixtralr=   0moonshine_streaming.modeling_moonshine_streamingr>   configuration_gemma4r@   rA   rB   rC   
get_loggerrd   loggerrl   rh   r^   ra   ro   rq   rs   r  rv   r   r   r   r  r  r$  Conv1dr6  r=  rK  rW  rq  r  r   r  r  r  r  r  r  r  r  r  r&  r5  rI  rK  rm  r  r  r`  r  r  r   r  r  r1  __all__rS   r_   r]   <module>rX     s      $ ! %   $ & ! . 3	 	 	 C S K F &  ^ ] E *  
 
 
 8 5 [ g g  
		H	%44<<4 LL4'4 T\	4
 ,,%4 4 
4nQ : Q*Q#@ Q2 
Q$; 
Q 
Q 
37 3  3BII :	N 	7ryy 7>i)299 i)X#bii #8; ;< RYY  H"bii "B&RYY &R0ryy 0l*3		 *3Z@0 @0Fai a 5&||5&	5& 
5& ,,	5&
 5& \\5&p:"6 :z :)O :) :)|"1 "J)H")) )H^^I ^U 5 UDq)")) q)h7 7"@ryy "@JO/ Od	$A 	=)2 =)@ `aW
o W
 bW
t ]^J
) J
 _J
ZSc, SclAH- AHH68 6>LL4'\\D( _:	u|| 	U\\ 	^c^j^j 	 p, ppf	 t%D ttnr_   