
    ^j"                       d dl Z d dlZd dlmZ d dlmZ d dlmZmZ d dl	Z	d dl
mZ d dlmc 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 ddlmZ ddlmZ ddl m!Z! ddl"m#Z#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/m0Z0m1Z1m2Z2 ddl3m4Z4 ddl5m6Z6m7Z7m8Z8m9Z9 ddl:m;Z; ddl<m=Z=m>Z> ddl?m@Z@mAZAmBZB  ed       G d dej                               ZD G d dej                        ZE G d d ej                        ZF G d! d"ej                        ZG G d# d$ej                        ZH G d% d&ej                        ZId' ZJd(e	j                  d)e	j                  d*e	j                  d+e	j                  d,eLe	j                  e	j                  f   f
d-ZMd.e	j                  d/eNd,e	j                  fd0ZO	 dWd1ej                  d2e	j                  d3e	j                  d4e	j                  d5e	j                  dz  d6ePd7ePd8e-e/   fd9ZQ G d: d;ej                        ZR G d< d=e!      ZS G d> d?ej                        ZTd@ ZUdXdAZV G dB dCej                        ZW G dD dEej                        ZX G dF dGe!      ZYe0e G dH dIe#                    ZZe0 G dJ dKe+             Z[ G dL dMe[      Z\e0 G dN dOe[             Z]e0 G dP dQe[             Z^e0e G dR dSe%                    Z_ G dT dUe[e      Z`g dVZay)Y    N)Callable)	dataclass)AnyOptional)	LayerNorm   )initialization)ACT2FN)CacheDynamicCache)GenerationMixin)use_kernel_forward_from_hub)create_causal_mask)FlashAttentionKwargs)GradientCheckpointingLayer)BaseModelOutputWithPastBaseModelOutputWithPoolingCausalLMOutputWithPast)ROPE_INIT_FUNCTIONSdynamic_rope_update)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)TransformersKwargsauto_docstringcan_return_tupletorch_compilable_check)deprecate_kwarg)accepts_precomputed_kwargsis_flash_attention_requestedmaybe_autocastmerge_with_config_defaults)capture_outputs)get_vision_cu_seqlensget_vision_position_ids   )Glm4vConfigGlm4vTextConfigGlm4vVisionConfigRMSNormc                   h     e Zd Zddeddf fdZdej                  dej                  fdZd Z xZ	S )	Glm4vRMSNormepsreturnNc                     t         |           t        j                  t	        j
                  |            | _        || _        y)z;
        Glm4vRMSNorm is equivalent to T5LayerNorm
        N)super__init__nn	Parametertorchonesweightvariance_epsilon)selfhidden_sizer-   	__class__s      s/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/glm4v/modeling_glm4v.pyr1   zGlm4vRMSNorm.__init__:   s1     	ll5::k#:; #    hidden_statesc                 "   |j                   }|j                  t        j                        }|j	                  d      j                  dd      }|t        j                  || j                  z         z  }| j                  |j                  |      z  S )N   T)keepdim)	dtypetor4   float32powmeanrsqrtr7   r6   )r8   r=   input_dtypevariances       r;   forwardzGlm4vRMSNorm.forwardB   sy    #))%((7 $$Q',,R,>%Ht?T?T4T(UU{{]--k:::r<   c                 ^    t        | j                  j                         d| j                   S )Nz, eps=)tupler6   shaper7   )r8   s    r;   
extra_reprzGlm4vRMSNorm.extra_reprI   s*    ))*+6$2G2G1HIIr<   )gư>)
__name__
__module____qualname__floatr1   r4   TensorrJ   rN   __classcell__r:   s   @r;   r,   r,   8   s7    $ $$ $;U\\ ;ell ;Jr<   r,   c                   ,     e Zd Zddef fdZd Z xZS )Glm4VisionMlpbiasc                    t         |           |j                  | _        |j                  | _        t        j                  | j                  | j                  |      | _        t        j                  | j                  | j                  |      | _        t        j                  | j                  | j                  |      | _	        t        |j                     | _        y NrX   )r0   r1   r9   out_hidden_sizeintermediate_sizer2   Linear	gate_projup_proj	down_projr
   
hidden_actact_fn)r8   configrX   r:   s      r;   r1   zGlm4VisionMlp.__init__N   s    !--!'!7!74#3#3T5K5KRVWyy!1!143I3IPTU4#9#94;K;KRVWV../r<   c                     | j                  | j                  | j                  |            | j                  |      z        S N)ra   rc   r_   r`   r8   hidden_states     r;   rJ   zGlm4VisionMlp.forwardW   s2    ~~dkk$..*FG$,,WcJddeer<   F)rO   rP   rQ   boolr1   rJ   rT   rU   s   @r;   rW   rW   M   s    0T 0fr<   rW   c                   `     e Zd Zdeddf fdZdej                  dej                  fdZ xZS )Glm4vVisionPatchEmbedrd   r.   Nc                 T   t         |           |j                  | _        |j                  | _        |j                  | _        |j
                  | _        | j                  | j                  | j                  g}t        j                  | j                  | j                  ||      | _	        y )N)kernel_sizestride)
r0   r1   
patch_sizetemporal_patch_sizein_channelsr9   	embed_dimr2   Conv3dproj)r8   rd   rn   r:   s      r;   r1   zGlm4vVisionPatchEmbed.__init__\   s     ++#)#=#= !--++//$//RIId..K`kl	r<   r=   c                 6   | j                   j                  j                  }|j                  d| j                  | j
                  | j                  | j                        }| j                  |j                  |            j                  d| j                        }|S )Nr@   rB   )	ru   r6   rB   viewrr   rq   rp   rC   rs   )r8   r=   target_dtypes      r;   rJ   zGlm4vVisionPatchEmbed.forwardf   s~    yy''--%**  $":":DOOT__
 		-"2"2"2"FGLLRQUQ_Q_`r<   	rO   rP   rQ   r)   r1   r4   rS   rJ   rT   rU   s   @r;   rl   rl   [   s5    m0 mT mU\\ ell r<   rl   c                        e Zd ZU ej                  ed<   d	dededdf fdZdej                  dej                  fdZ	 xZ
S )
Glm4vVisionRotaryEmbeddinginv_freqdimthetar.   Nc                     t         |           || _        || _        d|t	        j
                  d|dt        j                        |z  z  z  }| j                  d|d       y )N      ?r   r?   rw   r}   F
persistent)r0   r1   r~   r   r4   arangerR   register_buffer)r8   r~   r   r}   r:   s       r;   r1   z#Glm4vVisionRotaryEmbedding.__init__r   sY    
%ELLC%++$NQT$TUVZeDr<   position_idsc                 \    |j                  d      | j                  z  j                  d      S )Nr@   r&   )	unsqueezer}   flatten)r8   r   s     r;   rJ   z"Glm4vVisionRotaryEmbedding.forwardy   s'    &&r*T]]:CCAFFr<   )g     @)rO   rP   rQ   r4   rS   __annotations__intrR   r1   rJ   rT   rU   s   @r;   r|   r|   o   sI    llEC E ED EGELL GU\\ Gr<   r|   c                   n     e Zd Zd
dededededdf
 fdZdej                  dej                  fd	Z	 xZ
S )Glm4vVisionPatchMergerr~   context_dimrb   rX   r.   Nc                 x   t         |           t        j                  |||      | _        t        |      | _        t        j                  |||      | _        t        j                  |||      | _        t        j                  |||      | _	        t        j                         | _        t        |   | _        y rZ   )r0   r1   r2   r^   ru   r   post_projection_normr_   r`   ra   GELUact1r
   rc   )r8   r~   r   rb   rX   r:   s        r;   r1   zGlm4vVisionPatchMerger.__init__~   s    IIc3T2	$-cN!3$?yyk=;$?GGI	Z(r<   rh   c                     | j                  |      }| j                  | j                  |            }| j                  | j	                  | j                  |            | j                  |      z        S rf   )ru   r   r   ra   rc   r_   r`   rg   s     r;   rJ   zGlm4vVisionPatchMerger.forward   sY    yy.yy!:!:<!HI~~dkk$..*FG$,,WcJddeer<   ri   )rO   rP   rQ   r   strrj   r1   r4   rS   rJ   rT   rU   s   @r;   r   r   }   sJ    )C )c )s )$ )[_ )fELL fU\\ fr<   r   c                   D     e Zd Zdef fdZdej                  fdZ xZS )Glm4vVisionEmbeddingsrd   c                 f   t         |           || _        |j                  | _        |j
                  | _        |j                  | _        | j
                  | j                  z  dz  | _        | j                  | _        t        j                  | j                  | j                        | _        d| _        y )Nr?   bicubic)r0   r1   rd   r9   rs   
image_sizerp   num_patchesnum_positionsr2   	Embeddingposition_embeddinginterpolated_methodr8   rd   r:   s     r;   r1   zGlm4vVisionEmbeddings.__init__   s    ++ ++ ++ OOt>1D!--"$,,t/A/A4>>"R#, r<   r.   c                    | j                   j                  }|j                  d   }|j                  }t	        |t
              r&t        j                  ||t        j                        }|j                  d   }	t        |	dz        }
|j                  |
|
|      j                  ddd      j                  d      j                  |t        j                        }|j                  d   }t        j                  ||j                        }|j                  d      |j!                  d      j                  d      k\  j#                  d      }||df   j                  t        j                        }||df   j                  t        j                        }|dz   |z  dz  dz
  }|dz   |z  dz  dz
  }t        j$                  ||fd	      j                  d      j                  d      }t'        j(                  ||| j*                  d
d      }|j-                  d      j-                  d      j                  dd      }|j                  |j.                        j                  |j                        }||z   }|S )a  
        Forward pass with integrated position encoding adaptation using 2D interpolation.

        Args:
            embeddings: Input embeddings tensor
            lengths (torch.Tensor): Sequence lengths for each image in the batch.
            image_shapes (torch.Tensor): Tensor of shape [batch_size, 3] representing the image shapes (t, h, w).
            h_coords (torch.Tensor): Tensor of shape [total_seq] representing the h coordinate for each patch.
            w_coords (torch.Tensor): Tensor of shape [total_seq] representing the w coordinate for each patch.

        Returns:
            torch.Tensor: Embeddings with adapted position encoding added.
        r&   devicerB   r   g      ?r?   r   rw   r@   r~   Fborder)modealign_cornerspadding_mode)r   r6   rM   r   
isinstancelistr4   tensorlongr   rx   permuter   rC   rD   r   cumsumsumstackFgrid_sampler   squeezerB   )r8   
embeddingslengthsimage_shapesh_coordsw_coordspos_embed_weightr9   r   orig_size_sq	orig_sizepos_embed_2d
num_tokenstoken_positionsseq_idstarget_htarget_wnorm_wnorm_hgridinterpolated_embed_fp32adapted_pos_embed_fp32adapted_pos_embeds                          r;   rJ   zGlm4vVisionEmbeddings.forward   s=     2299&,,Q/!(( gt$ll76LG (--a0c)*	!!)YDWQ1Yq\RvU]]R3	 	  %%a(
,,z*:K:KL",,Q/7>>!3D3N3Nq3QQVVWXY
+..U]].C
+..U]].C c>X-2Q6c>X-2Q6 {{FF+4>>qAKKAN #$--$T%=%=Uai#

 "9!@!@!C!K!KB!O!W!WXY[\!]2556F6L6LMPPQ[QbQbc  "33
r<   rz   rU   s   @r;   r   r      s#    
-0 
-:PUP\P\ :r<   r   c                     | dd| j                   d   dz  f   }| d| j                   d   dz  df   }t        j                  | |fd      S )*Rotates half the hidden dims of the input..Nr@   r?   r   )rM   r4   catxx1x2s      r;   rotate_halfr      sZ    	
3"!''"+"""	#B	
3q ""	#B99rc2YB''r<   qkcossinr.   c                    | j                   }|j                   }| j                         |j                         }} |j                  d      j                         |j                  d      j                         }}| |z  t        |       |z  z   }||z  t        |      |z  z   }|j	                  |      }|j	                  |      }||fS )N)rB   rR   r   r   rC   )r   r   r   r   orig_q_dtypeorig_k_dtypeq_embedk_embeds           r;   apply_rotary_pos_emb_visionr      s     77L77L779aggiqA}}R &&(#--*;*A*A*CC3w;q>C/0G3w;q>C/0Gjj&Gjj&GGr<   r=   n_repc                     | j                   \  }}}}|dk(  r| S | dddddddddf   j                  |||||      } | j                  |||z  ||      S )z
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r&   N)rM   expandreshape)r=   r   batchnum_key_value_headsslenhead_dims         r;   	repeat_kvr      so    
 2?1D1D.Ehz!!Qa"23::5BUW\^bdlmM  (;e(CT8TTr<   modulequerykeyvalueattention_maskscalingdropoutkwargsc                    t        || j                        }t        || j                        }	t        j                  ||j	                  dd            |z  }
||
|z   }
t
        j                  j                  |
dt        j                        j                  |j                        }
t
        j                  j                  |
|| j                        }
t        j                  |
|	      }|j	                  dd      j                         }||
fS )Nr?   r   r@   )r~   rB   )ptrainingr&   )r   num_key_value_groupsr4   matmul	transposer2   
functionalsoftmaxrD   rC   rB   r   r   
contiguous)r   r   r   r   r   r   r   r   
key_statesvalue_statesattn_weightsattn_outputs               r;   eager_attention_forwardr      s     3 ; ;<JUF$?$?@L<<z';';Aq'ABWLL!#n4==((2U]](SVVW\WbWbcL==((6??([L,,|\:K''1-88:K$$r<   c            
            e Zd Zdeddf fdZ edd      	 ddej                  d	ej                  d
eej                  ej                  f   dz  dej                  fd       Z	 xZ
S )Glm4vVisionAttentionrd   r.   Nc                    t         |           |j                  | _        |j                  | _        | j                  | j                  z  | _        d| _        t        j                  |j                  |j                  dz  |j                        | _
        t        j                  |j                  |j                  d      | _        | j
                  dz  | _        || _        |j                  | _        d| _        y )Nr&   r   r[   F      )r0   r1   r9   r~   	num_headsr   r   r2   r^   attention_biasqkvru   r   rd   attention_dropout	is_causalr   s     r;   r1   zGlm4vVisionAttention.__init__  s    %%))DNN2$%!99V//1C1Ca1GfNcNcdIIf00&2D2D5Q	}}d*!'!9!9r<   rotary_pos_embv5.10versionr=   
cu_seqlensposition_embeddingsc                    |j                   d   }| j                  |      j                  |d| j                  d      j	                  dddd      j                  d      \  }}}|\  }	}
t        |||	|
      \  }}|j                  dd      j                  d      }|j                  dd      j                  d      }|j                  dd      j                  d      }t        j                  | j                  j                  t              }t        | j                        rT|dd  |d d z
  j                         } || |||fd | j                   | j"                  sdn| j$                  ||||dd|\  }}n|dd  |d d z
  }|||fD cg c](  }t'        j(                  ||j+                         d	      * }}t-        | D cg c]<  \  }}} || |||fd | j                   | j"                  sdn| j$                  dd
|d   > }}}}t'        j.                  |d	      }|j                  |d      j1                         }| j3                  |      }|S c c}w c c}}}w )Nr   r   r@   r&   r?           F)r   r   r   cu_seq_lens_qcu_seq_lens_kmax_length_qmax_length_kr  r   )r   r   r   r  )rM   r   r   r   r   unbindr   r   r   r   get_interfacerd   _attn_implementationr   r    maxr   r   r  r4   splittolistzipr   r   ru   )r8   r=   r  r  r   
seq_lengthquery_statesr   r   r   r   attention_interface
max_seqlenr   _r   r   splitsr   r   vattn_outputss                         r;   rJ   zGlm4vVisionAttention.forward   s    #((+
HH]#++J4>>2NVVWXZ[]^`abiijkl 	/j, 'S#>|ZY\^a#b j#--a3==a@))!Q/99!<
#--a3==a@(?(M(MKK,,.E)
 (4$QR.:cr?:??AJ0	
  $#'==d6L6L(('' NK" !nz#26GLXZdfrKsAGFGNN$4!<F    #F|  Aq! $	

 $( LL'+}}C$:P:P#
 
 
L   ))La8K!))*b9DDFii,-s   -I?AIrf   )rO   rP   rQ   r)   r1   r   r4   rS   rL   rJ   rT   rU   s   @r;   r   r     s    0 T  %w7
 IM	A||A LLA #5<<#=>E	A 
A 8Ar<   r   c                        e Zd Zd fdZ edd      e	 ddej                  dej                  d	eej                  ej                  f   dz  dej                  fd
              Z	 xZ
S )Glm4vVisionBlockr.   Nc                     t         |           t        |j                  |j                        | _        t        |j                  |j                        | _        t        |      | _        t        |d      | _
        y )Nr-   Fr[   )r0   r1   r,   r9   rms_norm_epsnorm1norm2r   attnrW   mlpr   s     r;   r1   zGlm4vVisionBlock.__init__f  s\    !&"4"4&:M:MN
!&"4"4&:M:MN
(0	 e4r<   r  r  r  r=   r  r  c                     | | j                   | j                  |      f||d|z   }|| j                  | j                  |            z   }|S )z
        cu_seqlens (`torch.Tensor`):
            Cumulative sequence lengths used for packed variable-length attention in Flash Attention kernels.
        r  r  )r%  r#  r&  r$  )r8   r=   r  r  r   s        r;   rJ   zGlm4vVisionBlock.forwardm  s`     &			JJ}%)
! 3)
 	)
 
 &M1J(KKr<   r.   Nrf   )rO   rP   rQ   r1   r   r   r4   rS   rL   rJ   rT   rU   s   @r;   r  r  e  s|    5 %w7
 IM	|| LL #5<<#=>E	 
  8r<   r  c                        e Zd ZU ej                  ed<   ddef fdZe	 	 	 ddedz  de	d   de
dz  ded	ef   fd
       Z ej                         ed               Zd Z xZS )Glm4vTextRotaryEmbeddingr}   Nrd   c                    t         |           |j                  | _        |j                  | _        || _        | j
                  j                  d   | _        | j                  }| j                  dk7  rt        | j                     } || j
                  |      \  }| _
        | j                  d|d       | j                  d|j                         d       |j                  j                  dg d      | _        y )	N	rope_typedefaultr}   Fr   original_inv_freqmrope_section)      r2  )r0   r1   max_position_embeddingsmax_seq_len_cachedoriginal_max_seq_lenrd   rope_parametersr-  compute_default_rope_parametersr   attention_scalingr   clonegetr0  )r8   rd   r   rope_init_fnr}   r:   s        r;   r1   z!Glm4vTextRotaryEmbedding.__init__  s    "("@"@$*$B$B!44[A!%!E!E>>Y&.t~~>L+7V+L($(ZeD0(..2BuU#3377Ur<   r   ztorch.deviceseq_lenr.   ztorch.Tensorc                 n   | j                   d   }| j                   j                  dd      }t        | dd      xs | j                  | j                  z  }t        ||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_thetapartial_rotary_factorr   r   Nr   r?   rw   r   )r6  r:  getattrr9   num_attention_headsr   r4   r   int64rC   rR   )	rd   r   r<  baser?  r   r~   attention_factorr}   s	            r;   r7  z8Glm4vTextRotaryEmbedding.compute_default_rope_parameters  s    & %%l3 & 6 6 : :;RTW X6:t4h8J8JfNhNh8h(223 U\\!S!5;;?BB&X]XcXcBdgjjk
 )))r<   c                 ^   | j                   d d d d d f   j                         j                  d|j                  d   dd      }|d d d d d d d f   j                         }t	        |j
                  j                  t              r/|j
                  j                  dk7  r|j
                  j                  nd}t        |d      5  |j                         |j                         z  j                  dd      }| j                  || j                        }t        j                  ||fd	      }|j                         | j                  z  }|j!                         | j                  z  }	d d d        j#                  |j$                  
      	j#                  |j$                  
      fS # 1 sw Y   AxY w)Nr   r&   r@   mpscpuF)device_typeenabledr?   r   rw   )r}   rR   r   rM   r   r   typer   r!   r   apply_mroper0  r4   r   r   r8  r   rC   rB   )
r8   r   r   inv_freq_expandedposition_ids_expandedrH  freqsembr   r   s
             r;   rJ   z Glm4vTextRotaryEmbedding.forward  s`   
 !MM$a*=>DDFMMaQ]QcQcdeQfhjlmn ,Q4] ; A A C'1!((--'E!((--[`J`ahhmmfkUC 	5&,,.1F1L1L1NNYYZ[]^_E$$UD,>,>?E))UEN3C'')d444C'')d444C	5 vvAGGv$cff177f&;;;	5 	5s   B!F##F,c           	          |}|j                  |d      }t        j                  t        |      D cg c]  \  }}||dz      c}}d      }|S c c}}w )Nr@   r   r   )r  r4   r   	enumerate)r8   rN  r0  sectionchunksichunkresults           r;   rK  z$Glm4vTextRotaryEmbedding.apply_mrope  sQ    W"-69JKXQE!a%LKQST Ls   A
rf   NNN)rO   rP   rQ   r4   rS   r   r(   r1   staticmethodr   r   rL   rR   r7  no_gradr   rJ   rK  rT   rU   s   @r;   r+  r+    s    llV V" )-+/"*$&*(* t* 
~u$	%	* *> U]]_<  < r<   r+  c                 |    | ddddf   }| ddddf   }t        j                  | |fd      j                  d      S )	r   .r   Nr?   r&   r@   r   r   )r4   r   r   r   s      r;   rotate_half_llmr[    sJ    	
319B	
319B;;Ryb)11"55r<   c                    |j                  |      }|j                  |      }|dd|j                  d   dz  f   j                  dd      }|dd|j                  d   dz  f   j                  dd      }|j                  d   }| dd|f   | d|df   }}|dd|f   |d|df   }	}||z  t        |      |z  z   }
||z  t        |      |z  z   }t	        j
                  |
|gd      }
t	        j
                  ||	gd      }|
|fS )a  Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        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:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    .Nr@   r?   r   )r   rM   repeat_interleaver[  r4   r   )r   r   r   r   unsqueeze_dim
rotary_dimq_rotq_passk_rotk_passr   r   s               r;   apply_rotary_pos_embrd    sD   $ --
&C
--
&C c'SYYr]a'''
(
:
:1"
:
EC
c'SYYr]a'''
(
:
:1"
:
EC 2Jc;J;&'3
+;)<6Ec;J;&'3
+;)<6E s{u5;<Gs{u5;<G ii&)r2Gii&)r2GGr<   c                   (    e Zd ZdZddededz  f fdZ	 	 	 ddej                  de	ej                  ej                  f   dz  dej                  dz  d	e
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 )Glm4vTextAttentionz
    Multi-headed attention from 'Attention Is All You Need' paper.
    and "Generating Long Sequences with Sparse Transformers".
    Nrd   	layer_idxc                    t         |           || _        || _        |j                  | _        |j
                  | _        | j                  | j                  z  | _        |j                  | _        | j                  | j                  z  | _	        d| _
        |j                  | _        |j                  | _        | j                  dz  | _        t        j                  | j                  | j                  | j                  z  d      | _        t        j                  | j                  | j                  | j                  z  d      | _        t        j                  | j                  | j                  | j                  z  d      | _        t        j                  | j                  | j                  z  | j                  d      | _        y )NTr   r[   F)r0   r1   rd   rg  r9   rA  r   r   r   r   r  r  r6  r   r2   r^   q_projk_projv_projo_projr8   rd   rg  r:   s      r;   r1   zGlm4vTextAttention.__init__  sI   "!--33((DNN:#)#=#= $(NNd6N6N$N!!'!9!9%55}}d*ii 0 0$..4==2PW[\ii 0 0$2J2JT]]2Zaefii 0 0$2J2JT]]2Zaefii >@P@PW\]r<   r=   r  r   past_key_valuesr   r.   c                 F   |j                         \  }}}| j                  |      }	| j                  |      }
| j                  |      }|	j	                  ||d| j
                        j                  dd      }	|
j	                  ||d| j
                        j                  dd      }
|j	                  ||d| j
                        j                  dd      }|\  }}t        |	|
||      \  }	}
| |j                  |
|| j                        \  }
}t        j                  | j                  j                  t              } || |	|
||f| j                  sdn| j                   | j"                  d|\  }}|j%                  ||d      j'                         }| j)                  |      }||fS )Nr@   r&   r?   r
  )r   r   )sizeri  rj  rk  rx   r   r   rd  updaterg  r   r  rd   r  r   r   r  r   r   r   rl  )r8   r=   r  r   rn  r   bszq_lenr  r  r   r   r   r   r  r   r   s                    r;   rJ   zGlm4vTextAttention.forward  s    &**,UA{{=1[[/
{{=1#((eRGQQRSUVW__S%T]]CMMaQRS
#((eRGQQRSUVW&S#7jRUWZ#[ j&'6'='=j,X\XfXf'g$J(?(M(MKK,,.E)
 %8	%
  $}}C$2H2HLL	%
 	%
!\ "))#ub9DDFkk+.L((r<   rf   rW  )rO   rP   rQ   __doc__r(   r   r1   r4   rS   rL   r   r   r   rJ   rT   rU   s   @r;   rf  rf     s    
^ ^3: ^. IM.2(,))||)) #5<<#=>E)) t+	))
 )) -.)) 
u||U\\D0%2E2LL	M))r<   rf  c                   V     e Zd Z fdZdej
                  dej
                  fdZ xZS )Glm4vTextMLPc                 *   t         |           || _        t        j                  |j
                  d|j                  z  d      | _        t        j                  |j                  |j
                  d      | _        t        |j                     | _        y )Nr?   Fr[   )r0   r1   rd   r2   r^   r9   r]   gate_up_projra   r
   rb   activation_fnr   s     r;   r1   zGlm4vTextMLP.__init__G  sp    IIf&8&8!f>V>V:V]bc6#;#;V=O=OV[\#F$5$56r<   r=   r.   c                     | j                  |      }|j                  dd      \  }}|| j                  |      z  }| j                  |      S )Nr?   r@   r   )rx  rU  ry  ra   )r8   r=   	up_statesgates       r;   rJ   zGlm4vTextMLP.forwardO  sL    %%m4	#//!/4i 2 24 88	~~i((r<   )rO   rP   rQ   r1   r4   FloatTensorrJ   rT   rU   s   @r;   rv  rv  F  s'    7)U%6%6 )5;L;L )r<   rv  c                   D    e Zd Zdedef fdZe	 	 	 	 	 ddej                  de	ej                  ej                  f   dz  dej                  dz  dej                  dz  d	edz  d
edz  de	ej                  e	ej                  ej                  f   dz  f   fd       Z xZS )Glm4vTextDecoderLayerrd   rg  c                    t         |           |j                  | _        t        ||      | _        t        |      | _        t        |j                  |j                        | _	        t        |j                  |j                        | _
        t        |j                  |j                        | _        t        |j                  |j                        | _        y )Nr!  )r0   r1   r9   rf  	self_attnrv  r&  r,   r"  input_layernormpost_attention_layernormpost_self_attn_layernormpost_mlp_layernormrm  s      r;   r1   zGlm4vTextDecoderLayer.__init__Y  s    !--+FI>'+F,>,>FDWDWX(4V5G5GVM`M`(a%(4V5G5GVM`M`(a%".v/A/AvGZGZ"[r<   Nr=   r  r   r   rn  	use_cacher.   c           
         |}| j                  |      } | j                  d||||||d|\  }}	| j                  |      }||z   }|}| j                  |      }| j	                  |      }| j                  |      }||z   }|S )N)r=   r  r   r   rn  r   )r  r  r  r  r&  r  )
r8   r=   r  r   r   rn  r  r   residualr  s
             r;   rJ   zGlm4vTextDecoderLayer.forwardc  s     !,,]; *4>> 
' 3)%+
 
q 55mD =0 !55mD///> =0r<   )NNNNF)rO   rP   rQ   r(   r   r1   r   r4   rS   rL   
LongTensorr   rj   r}  rJ   rT   rU   s   @r;   r  r  X  s    \ \3 \  IM.204(,!&#||# #5<<#=>E# t+	#
 &&-# # $;# 
u  %(9(95;L;L(L"MPT"TT	U# #r<   r  c                   :    e Zd ZU dZdZej                  dz  ed<   y)Glm4vModelOutputWithPast  
    rope_deltas (`torch.LongTensor` of shape `(batch_size, )`, *optional*):
        The rope index difference between sequence length and multimodal rope.
        The attribute is deprecated and will be removed in v5.20, use `model.base_model.rope_deltas` instead.
    Nrope_deltasrO   rP   rQ   rt  r  r4   r  r   r  r<   r;   r  r         ,0K!!D(/r<   r  c                   T     e Zd ZU eed<   dZdZdZddgZdgZ	dZ
dZdZdZ fdZ xZS )	Glm4vPreTrainedModelrd   model)imagevideotextTr  r  rn  c                 "   t         |   |       t        |t              rod|j                  t        j                  d|j                  dt
        j                        |j                  z  z  z  }t        j                  |j                  |       y y )Nr   r   r?   rw   )r0   _init_weightsr   r|   r   r4   r   r~   rR   initcopy_r}   )r8   r   r}   r:   s      r;   r  z"Glm4vPreTrainedModel._init_weights  sk    f%f89fllu||Avzz1TYT_T_/`cicmcm/mnoHJJv1 :r<   )rO   rP   rQ   r'   r   base_model_prefixinput_modalitiessupports_gradient_checkpointing_no_split_modules_skip_keys_device_placement_supports_flash_attn_supports_sdpa_can_compile_fullgraph_supports_attention_backendr  rT   rU   s   @r;   r  r    sQ    1&*#02DE#4"5N!"&2 2r<   r  c                        e Zd ZU eed<   dZdgZeedZ	d fdZ
d Zeeedej                   d	ej                   d
ee   deez  fd                     Z xZS )Glm4vVisionModelrd   )r  r  r  r=   
attentionsr.   c                 F   t         |   |       |j                  | _        |j                  | _        t	        |      | _        t        |      | _        |j                  |j                  z  }t        |dz        | _        t        j                  t        |j                        D cg c]  }t!        |       c}      | _        t%        |j&                  |j(                  |j*                        | _        t/        |j                  |j0                        | _        t        j4                  |j                  |j&                  |j                  |j                        | _        t/        |j                  |j0                        | _        d| _        | j=                          y c c}w )Nr?   )r~   r   rb   r!  )rr   out_channelsrn   ro   F)r0   r1   spatial_merge_sizerp   r   r   rl   patch_embedr9   r   r|   r  r2   
ModuleListrangedepthr  blocksr   r\   r]   rb   mergerr,   r"  post_conv_layernormConv2d
downsamplepost_layernormgradient_checkpointing	post_init)r8   rd   r   r  r:   s       r;   r1   zGlm4vVisionModel.__init__  sA    "(";"; ++/708%%)9)998QGmmuV\\GZ$[!%5f%=$[\,&&F4L4LY_YjYj
 $00B0BH[H[#\ ))**//11,,	
 +6+=+=6CVCVW&+# %\s   %Fc                     t        j                  d| j                  j                   dt        d       t        || j                        }| j                  |      }||fS )N`z.rot_pos_emb` is deprecated and will be removed in v5.11. Use `get_vision_position_ids` from `transformers.vision_utils` and apply the rotary embedding module.r?   )
stacklevel)warningswarnr:   rO   FutureWarningr%   r  r  )r8   grid_thwr   r  s       r;   rot_pos_embzGlm4vVisionModel.rot_pos_emb  sa    ''(  )H  I	

 /x9P9PQ,,\:|++r<   r=   r  r   c           	      x   t        || j                  |      }t        ||      }| j                  |      }| j	                  |      }| j                  |      }t        j                  ||fd      }|j                         |j                         f}|dd |dd z
  }	| j                  ||	||dddf   j                  |j                        |dddf   j                  |j                              }| j                  D ]  }
 |
|f||d|} | j                  |      }|j                  d| j                  | j                  |j                   d         }|j#                  dddd	      }| j%                  |      j                  d| j&                  j(                        }| j+                  |      }t-        ||
      S )a\  
        hidden_states (`torch.Tensor` of shape `(seq_len, hidden_size)`):
            The final hidden states of the model.
        grid_thw (`torch.Tensor` of shape `(num_images_or_videos, 3)`):
            The temporal, height and width of feature shape of each image in LLM.

        Returns:
            `torch.Tensor`: hidden_states.
        )r   r@   r   r&   Nr   r(  r   r?   )last_hidden_statepooler_output)r%   r  r$   r  r  r  r4   r   r   r   r   rC   r   r  r  rx   rM   r   r  rd   r\   r  r   )r8   r=   r  r   r   r  
rotary_embrO  r  seqlensblkmerged_hidden_statess               r;   rJ   zGlm4vVisionModel.forward  s    /x9P9PY_`*8FC
((700?((6
iiZ0b9"wwy#'')4QR.:cr?2A!!-"6"67A!!-"6"67
 ;; 	C%$7 	M	 ++M:%**'')@)@-BUBUVXBY
 &--aAq96;;B@[@[\#{{=9)+.
 	
r<   r)  )rO   rP   rQ   r)   r   r  r  r  r   _can_record_outputsr1   r  r"   r#   r   r4   rS   r   r   rL   r   rJ   rT   rU   s   @r;   r  r    s    )+,)*
8,  3
"\\3
5:\\3
MSTfMg3
	+	+3
    3
r<   r  c                       e Zd ZU eed<   dZeedZde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dz  d
ej"                  dz  dedz  dee   deez  fd                     Z xZS )Glm4vTextModelrd   )r  r  c           	         t         |   |       |j                  | _        |j                  | _        t        j                  |j                  |j                  | j                        | _        t        j                  t        |j                        D cg c]  }t        ||       c}      | _        t        |j                  |j                        | _        t#        |      | _        d| _        | j)                          y c c}w )Nr!  rd   F)r0   r1   pad_token_idpadding_idx
vocab_sizer2   r   r9   embed_tokensr  r  num_hidden_layersr  layersr,   r"  normr+  r  r  r  rm  s      r;   r1   zGlm4vTextModel.__init__  s     !.. ++LL):):F<N<NPTP`P`ammGLVMeMeGfg)"695g
 !!3!39L9LM	2&A&+# hs   DN	input_idsr   r   rn  inputs_embedsr  r   r.   c           	      T   |d u |d uz  rt        d      |r6|4t        j                  j                         st	        | j
                        }|| j                  |      }|w||j                         nd}t        j                  |j                  d   |j                        |z   }|j                  ddd      j                  d|j                  d   d      }n2|j                  dk(  r#|d	   j                  d|j                  d   d      }|j                  dk(  r|j                  d   d
k(  r|d   }	|dd  }nd }	| j
                  ||||	d}
t        di |
}|}| j                  ||      }| j                   D ]  } ||f||	||d|}|} | j#                  |      }t%        ||      S )N:You must specify exactly one of input_ids or inputs_embedsr  r   r&   r   r@   r   r?   N.   )rd   r  r   rn  r   )r   )r   r   rn  r  )r  rn  r  )
ValueErrorr4   jit
is_tracingr   rd   r  get_seq_lengthr   rM   r   rx   r   ndimr   r  r  r  r   )r8   r  r   r   rn  r  r  r   past_seen_tokenstext_position_idsmask_kwargscausal_maskr=   r  decoder_layerlayer_outputss                   r;   rJ   zGlm4vTextModel.forward,  s    -t";<YZZ 09M9M9O*$++>O  --i8M CRC^==?de <<(;(;A(>}G[G[\_ooL',,Q26==aATATUVAWY[\L!#'	299!\=O=OPQ=RTVWL !l&8&8&;q&@ ,Q'+L !% kk*,.-
 )7;7%"oom,oW![[ 		*M)*. /$7 M *M		* 		-0&++
 	
r<   )NNNNNN)rO   rP   rQ   r(   r   r  r  rf  r  r1   r   r"   r#   r4   r  rS   r   r}  rj   r   r   rL   r   rJ   rT   rU   s   @r;   r  r    s     .(
    .2.204(,26!%J
##d*J
 t+J
 &&-	J

 J
 ((4/J
 $;J
 -.J
 
(	(J
    J
r<   r  c                   h    e Zd ZdZdZddgZ fdZ	 	 	 	 d(dedeeeef   e	j                  z  d	ed
ededee	j                  z  dz  fdZ	 	 	 d)de	j                  de	j                  de	j                  dz  de	j                  dz  de	j                  dz  dee	j                  e	j                  f   fdZ ed      ee	 d*de	j*                  de	j                  dz  dee   deez  fd                     Z ed      ee	 d*de	j*                  de	j                  dz  dee   deez  fd                     Z	 	 d+de	j                  de	j*                  de	j*                  dz  de	j*                  dz  f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  fd"Z ed#d$%      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	j*                  dz  de	j                  dz  de	j                  dz  de	j                  dz  dee   deez  fd'                     Z  xZ!S ).
Glm4vModelr  Fr  r  c                     t         |   |       t        j                  |j                        | _        t        j                  |j                        | _        d | _	        | j                          y rf   )r0   r1   r  _from_configvision_configvisualr  text_configlanguage_modelr  r  r   s     r;   r1   zGlm4vModel.__init__  sU     &33F4H4HI,99&:L:LM 	r<   Nstart_positionr  temp_merge_sizer  time_intervalr   c                    |d   j                         |z  |d   j                         |z  |d   j                         |z  }	}}t        j                  ||      |z  }
t        j                  ||      |z   }t        j                  |	|      |z   }t        j                  |
||d      \  }}}t        j                  |||gd      j                  dd	      }|dxx   |z  cc<   |S )
a  
        Compute 3D positional indices for vision tokens derived from a single image or video input.

        The positions are generated from the input grid defined by temporal (T), height (H), and
        width (W) dimensions. Temporal and spatial dimensions can be downscaled according to the
        merge sizes used in the vision backbone. The resulting positions are offset by `start_position`.

        Args:
            start_position (`int`):
                Offset added to all computed positional indices.
            grid_thw (`Sequence[int]` or `torch.Tensor` of shape `(3,)`):
                The (T, H, W) grid representing the feature layout of the current image or video after patch embedding.
            temp_merge_size (`int`, *optional*):
                Factor by which the temporal dimension is reduced in the backbone. The temporal grid size is divided
                by this value. Defaults to 1.
            spatial_merge_size (`int`, *optional*):
                Factor by which the spatial dimensions (H and W) are reduced in the backbone. Both H and W are divided
                by this value. Defaults to 1.
            time_interval (`int`, *optional*):
                Spacing factor applied between consecutive temporal position indices.Defaults to 1.
            device (`str` or `torch.device`, *optional*):
                Device on which the resulting tensor is allocated. If `None`, uses the current default device.

        Returns:
            torch.LongTensor of shape (3, sequence_length):
                Positional indices for temporal, height, and width dimensions,
                flattened into sequence form and offset by `start_position`.
        r   r&   r?   r   ij)indexingr   r   r@   )itemr4   r   meshgridr   r   )r8   r  r  r  r  r  r   
llm_grid_t
llm_grid_h
llm_grid_wposition_temporalposition_heightposition_widthT_gridH_gridW_gridvision_position_idss                    r;   r%   z"Glm4vModel.get_vision_position_ids  s    L QK/1QK"44QK"44 !+J
 "LLFCmS,,z&ANRj@>Q!&0A?Tbmq!r#kk666*BJRRSTVXYA.0""r<   r  mm_token_type_idsimage_grid_thwvideo_grid_thwr   r.   c           	      "   |(t        j                  ||dddf   d      }d|dddf<   | j                  j                  j                  }g }t        j
                  d|j                  d   |j                  d   |j                  |j                        }	|t        |      nd|t        |      ndd}
t        |      D ]  \  }}||   }|,|||   j                            }|||   j                            }g }t        j                  t        |j                               d       D ]7  \  }}t        |      }|d   d   }|d	   d   dz   }|j!                  |||f       9 d}g }|D ]  \  }}}|dk(  r^||z
  }|j!                  t        j"                  ||j                  
      j%                  dd	      j'                  dd	      |z          ||z  }jt)        |
|         }| j+                  ||d||j                  
      }|j!                  |       |t-        |d   |d         |z  z  } t        j.                  |d      j1                  dd	      }|5|j3                  |	j                        |	dd|||   j                         f<   n"|j3                  |	j                        |	dd|f<   |j!                  |j-                         dz   t5        |      z
          t        j6                  ||j                  
      j9                  d      }|	|fS )a  
        Difference from Qwen2VL/Qwen2.5VL's get_rope_index:
        - GLM4V uses timestamps to separate each video frame, so the video_grid_thw should also be split too.

        Args:
            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
                Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
                it.
            mm_token_type_ids (`torch.IntTensor` of shape `(batch_size, sequence_length)`):
                Token type ids matching each modality to a different value in the input sequence, i.e. text (0), image (1), video (2).
            image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
                The temporal, height and width of feature shape of each image in LLM.
            video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
                The temporal, height and width of feature shape of each video in LLM.
            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

                - 1 for tokens that are **not masked**,
                - 0 for tokens that are **masked**.

        Returns:
            position_ids (`torch.LongTensor` of shape `(3, batch_size, sequence_length)`)
            mrope_position_deltas (`torch.Tensor` of shape `(batch_size)`)
        Nr   r   r&   r   rB   r   )r&   r?   c                     | d   S )Nr&   r  )r   s    r;   <lambda>z+Glm4vModel.get_rope_index.<locals>.<lambda>  s    `abc`d r<   r@   r   r?   )r4   r]  rd   r  r  zerosrM   rB   r   iterrQ  rj   	itertoolsgroupbyr  r   appendr   rx   r   nextr%   r  r   r   rC   lenr   r   )r8   r  r  r  r  r   r   r  mrope_position_deltasr   
grid_iters	batch_idxcurrent_input_idsinput_token_typeinput_type_groupr   groupstart_index	end_indexcurrent_posllm_pos_ids_listmodality_type	start_idxend_idxtext_lenr  r  llm_positionss                               r;   get_rope_indexzGlm4vModel.get_rope_index  sO   F %"44^^TUWXTXEY_`aN#$N1a4 ![[66II "{{OOAOOA//##
 (6'AtN#t'5'AtN#t


 -6i,@ $	[(I(0;)$5nY6O6T6T6V$W!#3N94M4R4R4T#U !'//	:J:Q:Q:S0TVde G
UU#Ahqk!"IaL1,	 ''k9(EF	G K!5E W1y' A%&2H$++Xi6F6FGLLQPRSZZ[\^`adoo  8+K  $J}$=>H*.*F*F#Xq2DYM]M] +G +' %++,?@3x{HQK#@DV#VVKW  "II&6A>FFq"MM)O\O_O_`l`s`sOtQ	>)+D+I+I+KKL-:-=-=l>Q>Q-RQ	\*!(():):)<q)@3GXCY)YZI$	[J !&-B9K[K[ \ f fgh i222r<   r  )modalitypixel_values_videosr   c                    |j                  | j                  j                        }|dddf   }|ddddf   }t        j                  ||d      }|j                  |j                  d   d      }t        j                  ||gd      } | j                  |f|dd|}	|j                  d      | j                  j                  dz  z  j                         }
t        j                  |	j                  |
      }||	_        |	S )	[  
        pixel_values_videos (`torch.FloatTensor` of shape `(batch_size, num_channels, image_size, image_size)`):
            The tensors corresponding to the input videos.
        video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
            The temporal, height and width of feature shape of each video in LLM.
        Nr   r&   r   T)r  return_dictr@   r?   )rJ  r  rB   r4   r]  new_onesrM   r   prodr  r  r  r  )r8   r  r  r   thwflattened_hwprefix_onesflattened_video_grid_thwvision_outputssplit_sizesvideo_embedss               r;   get_video_featureszGlm4vModel.get_video_features  s     266t{{7H7HI1a4 AqrE"..r1!<$--l.@.@.CQG#(99k<-Ha#P $
*BPT
X^
 &**2.$++2P2PRS2SS[[]{{>#?#?M'3$r<   r  pixel_valuesc                 :   |j                  | j                  j                        } | j                  |fd|i|}|j                  d      | j                  j                  dz  z  j                         }t        j                  |j                  |      }||_        |S )T  
        pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, image_size, image_size)`):
            The tensors corresponding to the input images.
        image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
            The temporal, height and width of feature shape of each image in LLM.
        r  r@   r?   )	rJ  r  rB   r"  r  r  r4   r  r  )r8   r,  r  r   r(  r)  image_embedss          r;   get_image_featureszGlm4vModel.get_image_features=  s     $(():):;$\UNUfU%**2.$++2P2PRS2SS[[]{{>#?#?M'3$r<   r  image_featuresvideo_featuresc                    || | j                         t        j                  | j                  j                  t        j
                  |j                              k(  }|j                  d      }| | j                         t        j                  | j                  j                  t        j
                  |j                              k(  }|j                  d      }n2|| j                  j                  k(  }|| j                  j                  k(  }|j                         }|j                  d      j                  |j                        }|@t        ||j                  d   z  |j                         k(  d| d|j                  d           |j                         }|j                  d      j                  |j                        }|@t        ||j                  d   z  |j                         k(  d| d|j                  d           ||fS )z
        Obtains multimodal placeholder mask from `input_ids` or `inputs_embeds`, and checks that the placeholder token count is
        equal to the length of multimodal features. If the lengths are different, an error is raised.
        r  r@   z6Image features and image tokens do not match, tokens: z, features: r   z6Video features and video tokens do not match, tokens: )get_input_embeddingsr4   r   rd   image_token_idr   r   allvideo_token_idr   r   rC   r   rM   numel)	r8   r  r  r1  r2  special_image_maskspecial_video_maskn_image_tokensn_video_tokenss	            r;   get_placeholder_maskzGlm4vModel.get_placeholder_maskT  s    !.2M$2K2K2MT[[77uzzR_RfRfg3 " "4!7!7!;!.2M$2K2K2MT[[77uzzR_RfRfg3 " "4!7!7!; "+dkk.H.H!H!*dkk.H.H!H+//1/99"=@@AUAUV%"!4!4R!88N<P<P<RRHHXXdeseyeyz{e|d}~
 ,//1/99"=@@AUAUV%"!4!4R!88N<P<P<RRHHXXdeseyeyz{e|d}~ "#555r<   rn  c                    |dn|j                         }|d uxs |d u}	|	r||t        d      |d uxr |d uxr |	}
|
r3| j                  |dk(  r"| j                  |||||      \  }}|| _        |S | j                  =|dkD  s|5|j                  \  }}}|u|j                         j                  d      dz
  }|j                  |dk(  d      }|j                  d|d      j                  ddd      j                  |j                        }nVt        j                  |||z         }|j                  ddd      j                  d|d      j                  |j                        }| j                  j                  || j                  j                  d   z  d      }||j                  |j                        z   }|S d }|S )	Nr   a  Multimodal data was passed (via `image_grid_thw` or `video_grid_thw`) but `mm_token_type_ids` is missing. Please pass `mm_token_type_ids` to the model so that multimodal RoPE (M-RoPE) can be computed correctly. `mm_token_type_ids` is returned by the processor alongside `input_ids`.)r  r  r   r  r@   r&   r   r   r   )r  r  r  r  rM   r   r   masked_fillrx   repeatrC   r   r4   r   r   r]  )r8   r  r  r  r  r   rn  r  past_key_values_lengthhas_multimodalcan_compute_mroper   r  
batch_sizer  r  deltas                    r;   compute_3d_position_idsz"Glm4vModel.compute_3d_position_ids~  s    '6&=?CaCaCc't3Q~T7Q/7I<Qn 
 &T1f6Gt6SfXf$"2"2":>TXY>Y(,(;(;---"3 )< )%L+  +D&  )/E/IYM^(5(;(;%J
A)-224;;B?!C+77!8KQO+00JCJJ1aQRSVVWdWkWkl$||,BDZ]gDgh+00Ar:AA!ZQSTWWXeXlXlm$$66zTEUEUE[E[\]E^7^de6fE'%((-:N:N(*OOL   Lr<   r  r  r  r   c           	         |du |duz  rt        d      | | j                         |      }| | j                  ||fddi|j                  }t	        j
                  |d      j                  |j                  |j                        }| j                  |||      \  }}|j                  ||      }| | j                  ||	fddi|j                  }t	        j
                  |d      j                  |j                  |j                        }| j                  |||      \  }}|j                  ||      }|| j                  |||	||||
	      } | j                  dd||||d
|}t        di |d| j                  iS )aU  
        image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
            The temporal, height and width of feature shape of each image in LLM.
        video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
            The temporal, height and width of feature shape of each video in LLM.
        Nr  r   Tr   r   )r1  )r2  )r  r  r  r  r   rn  r  )r  r   r   rn  r  r  r  )r  r4  r0  r  r4   r   rC   r   rB   r=  masked_scatterr+  rF  r  r  r  )r8   r  r   r   rn  r  r,  r  r  r  r  r   r/  
image_maskr  r*  
video_maskoutputss                     r;   rJ   zGlm4vModel.forward  s   . -t";<YZZ 7D557	BM#2422n:>BHm  !99\q9<<]=Q=QS`SfSfgL 55i_k5lMJ)88\RM*2422#^AEIOm  !99\q9<<]=Q=QS`SfSfgL 55i_k5lMAz)88\RM77#--+- /"3 8 L &$%% 
%)+'
 
 ( 

((
 	
r<   )r&   r&   r&   NrW  rf   )NN)NNNNN)
NNNNNNNNNN)"rO   rP   rQ   r  accepts_loss_kwargsr  r1   r   r   r4   rS   r   r   r%   r  	IntTensorrL   r  r   r   r   r}  r   r   r   r+  r0  r=  rF  r   r   r  rJ   rT   rU   s   @r;   r  r  |  s>   02DE  !"#,02#2# sC}%42# 	2#
  2# 2# ell"T)2#p 3726.2[3##[3 !??[3 ((4/	[3
 ((4/[3 t+[3 
u||U\\)	*[3z  1 37".. ((4/ +,	
 
+	+   2:  1 37'' ((4/ +,	
 
+	+   20 4837(6##(6 (((6 ))D0	(6
 ))D0(6\ /3.2.2/348/<<$&/ ||d*/ t+	/
 t+/ t+/ ,/ !??T1/ 
	/b ]G4 .2.204(,26,08<262648A
##d*A
 t+A
 &&-	A

 A
 ((4/A
 llT)A
 #..5A
 ((4/A
 ((4/A
 !??T1A
 +,A
 
)	)A
   5A
r<   r  c                   :    e Zd ZU dZdZej                  dz  ed<   y)Glm4vCausalLMOutputWithPastr  Nr  r  r  r<   r;   rO  rO    r  r<   rO  c            !           e Zd ZddiZdZ fdZe	 d dej                  dej                  dz  de
e   d	eez  fd
       Ze	 d dej                  dej                  dz  de
e   d	eez  fd       Z edd      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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ej$                  z  de
e   d	eez  fd                     Z	 	 	 	 	 	 	 	 	 	 d" fd	Z fdZ	 d dej                  dz  dej$                  dz  d	eej$                  ej$                  f   fdZ	 	 	 d#dededej                  dz  d	eej                  eeef   f   fdZ xZ S )$Glm4vForConditionalGenerationzlm_head.weightz(model.language_model.embed_tokens.weightFc                     t         |   |       t        |      | _        t	        j
                  |j                  j                  |j                  j                  d      | _	        | j                          y )NFr[   )r0   r1   r  r  r2   r^   r  r9   r  lm_headr  r   s     r;   r1   z&Glm4vForConditionalGeneration.__init__  sS     '
yy!3!3!?!?ASASA^A^ejkr<   Nr  r  r   r.   c                 >     | j                   j                  ||fi |S )r  )r  r+  )r8   r  r  r   s       r;   r+  z0Glm4vForConditionalGeneration.get_video_features  s$     -tzz,,-@.[TZ[[r<   r,  r  c                 >     | j                   j                  ||fi |S )r.  )r  r0  )r8   r,  r  r   s       r;   r0  z0Glm4vForConditionalGeneration.get_image_features  s"     -tzz,,\>TVTTr<   r  r  r  r  r   r   rn  r  labelsr  logits_to_keepc                     | j                   d||||	|
|||||d
|}|d   }t        |t              rt        | d      n|}| j	                  |dd|ddf         }d}|2| j                  ||| j                  j                  j                        }t        |||j                  |j                  |j                  |j                        S )a  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
        image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
            The temporal, height and width of feature shape of each image in LLM.
        video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
            The temporal, height and width of feature shape of each video in LLM.

        Example:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoProcessor, Glm4vForConditionalGeneration

        >>> model = Glm4vForConditionalGeneration.from_pretrained("zai-org/GLM-4.1V-9B-Thinking")
        >>> processor = AutoProcessor.from_pretrained("zai-org/GLM-4.1V-9B-Thinking")

        >>> messages = [
            {
                "role": "user",
                "content": [
                    {"type": "image", "url": "https://www.ilankelman.org/stopsigns/australia.jpg"},
                    {"type": "text", "text": "What is shown in this image?"},
                ],
            },
        ]
        >>> url = "https://www.ilankelman.org/stopsigns/australia.jpg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
        >>> inputs = processor(text=[text], images=[image], vision_infos=[vision_infos])

        >>> # 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]
        "The image shows a street scene with a red stop sign in the foreground. In the background, there is a large red gate with Chinese characters ..."
        ```)
r  r,  r  r  r  r  r   r   rn  r  r   N)logitsrV  r  )lossrY  rn  r=   r  r  r  )r  r   r   slicerS  loss_functionrd   r  r  rO  rn  r=   r  r  )r8   r  r   r   rn  r  rV  r,  r  r  r  r  rW  r   rK  r=   slice_indicesrY  rZ  s                      r;   rJ   z%Glm4vForConditionalGeneration.forward,  s    z $** 
% 3))/%)+'
 
  
 9C>SV8W~ot4]kmA}a,?@A%%VFt{{OfOfOqOq%rD*#33!//))++
 	
r<   c                 Z    t        |   |f|||||||	|
||d
|}|s|r
d |d<   d |d<   |S )N)
rn  r   r  r   r,  r  r  r  r  is_first_iterationr,  r  )r0   prepare_inputs_for_generation)r8   r  rn  r   r  r   r  r,  r  r  r  r_  r   model_inputsr:   s                 r;   r`  z;Glm4vForConditionalGeneration.prepare_inputs_for_generation  sf    " w<
+)'%% 3))1
 
 "i+/L(26L./r<   c                    t         |   ||      }d}|j                  d      x}|j                         }|dk7  r4| j                  j
                  |d   | j                  j
                  z   }|S d|v r|d   j                  d   dkD  r|d   }t        |j                        dk(  xr, |j                  t        j                  t        j                  fv }|r|j                  d      }|j                  d      |j                  d	      [|j                         D 	ci c]  \  }}	|dk7  s||	 }}}	 | j                  j                  |fi |\  }
}|| j                  _        no|j                  d      j                  d
dd      }
t        j                   |j                  d   dt        j                  |j"                        | j                  _        |d   }t        j$                  ||
gd      }|S c c}	}w )Nr   rn  r  r  r&   r?   r  r  r  r   r@   r  r   )r0   $_prepare_position_ids_for_generationr:  r  r  r  rM   r
  rB   r4   r   r   itemsr  r   r   r  r   r   )r8   inputs_tensormodel_kwargstext_positionspast_lengthcacher   is_input_idsr   r  vision_positionsr  r:   s               r;   rc  zBGlm4vForConditionalGeneration._prepare_position_ids_for_generation  s    EmUab !%%&788EE..0K!

 6 6 B))4tzz7M7MML ,&<+D+J+J1+MPQ+Q(5M=../14g9L9LQVQZQZ\a\f\fPg9g  !45A!!"23?<CSCSTdCeCq-9-?-?-AVTQQ+EUAqDVLV,EDJJ,E,Em,dWc,d)k%0DJJ"-77:AA!RL%*[[##A&MDXDX&DJJ"
 (	2yy.2B!CK Ws   G3*G3c                    || | j                         t        j                  | j                  j                  t        j
                  |j                              k(  d   }| | j                         t        j                  | j                  j                  t        j
                  |j                              k(  d   }| | j                         t        j                  | j                  j                  t        j
                  |j                              k(  d   }nK|| j                  j                  k(  }|| j                  j                  k(  }|| j                  j                  k(  }t        j                  |j                         |j                         z
  d      }|dkD  }|| z  }|j                  d      }	|j                  d      }
|	|
fS )aa  
        Get the number of images and videos for each sample to calculate the separation length of the sample tensor.
        These parameters are not passed through the processor to avoid unpredictable impacts from interface modifications.

        Args:
            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
                Indices of input sequence tokens in the vocabulary.

        Returns:
            image_nums (`torch.LongTensor` of shape `(batch_size, num_images_sample)`)
            video_nums (`torch.LongTensor` of shape `(batch_size, num_videos_sample)`)
        r  ).r   r&   r   r   )r4  r4   r   rd   image_start_token_idr   r   video_start_token_idvideo_end_token_idr   r   r   )r8   r  r  is_imageis_video_startis_video_endvideo_levelinside_videostandalone_imagesimage_countsvideo_countss              r;   _get_image_nums_and_video_numsz<Glm4vForConditionalGeneration._get_image_nums_and_video_nums  s   $ $.4,,.LL!A!A\i\p\pq H .4,,.LL!A!A\i\p\pq N .4,,.LL!?!?uzzZgZnZno L !DKK$D$DDH&$++*J*JJN$(F(FFL ll>#5#5#7,:J:J:L#LRST"Q %6 ),,,3%))a)0\))r<   expand_sizeis_encoder_decoderc                      dk(  rfS g d fd}fd} |      j                  d       |      |r*j                  d      t        d       |d         d<   fS )	Nr&   )r,  r  r  r  second_per_grid_tsc                 4   j                  dd       }j                  dd       }j                  j                  dd             \  }}d }| D ]9  }|dk(  rct        j                  |t	        |            }|D cg c]'  }t        j
                  |d      j                         ) }	} || |   |	
	      | |<   l|dk(  rt	        |      }	 || |   |	
	      | |<   |d
k(  rct        j                  |t	        |            }|D cg c]'  }t        j
                  |d      j                         ) }	} || |   |	
	      | |<   |dk(  rt	        |      }	 || |   |	
	      | |<   |dk(  s  || |   t	        |      
	      | |<   < | S c c}w c c}w )Nr  r  r  )r  c                     t        j                  | |      }|gdg| j                         dz
  z  z   }t        j                  |D cg c]  } |j                  |  c}d      }|S c c}w )Nr&   r   r   )r4   r  r~   r   r@  )r   r   repeat_timessamplesrepeat_argssamplerV  s          r;   _repeat_interleave_sampleszGlm4vForConditionalGeneration._expand_inputs_for_generation.<locals>._expand_dict_for_generation_visual.<locals>._repeat_interleave_samples&  sa    ++a1+nsaeegk/BBg#VFMFMM;$?#V\]^ $Ws   A&r,  r&   r   )r   r  r  r|  )r:  rx  r4   r  r   r"  r   )dict_to_expandr  r  
image_nums
video_numsr  r   r  r  r   ry  r  rf  r8   s             r;   "_expand_dict_for_generation_visualzgGlm4vForConditionalGeneration._expand_inputs_for_generation.<locals>._expand_dict_for_generation_visual  s   )--.>EN)--.>EN%)%H%H)9)9/4)P &I &"J
 & .(#kk.$z:JKGMTU6uzz&a8<<>UGU*D&s+W;+N3' ,,":.G*D&s+W;+N3' 11#kk.$z:JKGMTU6uzz&a8<<>UGU*D&s+W;+N3' ,,":.G*D&s+W;+N3' 00*D&s+T*5ET_+N3'7< "!3 V Vs   =,F,Fc                     | D ]u  }|dk(  r,| |   j                   dk(  r| |   j                  d      | |<   4| |   :t        | |   t        j                        sX|vs]| |   j                  d      | |<   w | S )Nr   r   r&   r   r   )r  r]  r   r4   rS   )r  r   ry  visual_keyss     r;   _expand_dict_for_generationz`Glm4vForConditionalGeneration._expand_inputs_for_generation.<locals>._expand_dict_for_generationL  s    % d.(^C-@-E-E-J*8*=*O*OP[ab*O*cN3'"3'3">##6E;.*8*=*O*OP[ab*O*cN3'd "!r<   r   r   encoder_outputszMIf `is_encoder_decoder` is True, make sure that `encoder_outputs` is defined.)r]  r:  r  )r8   ry  rz  r  rf  r  r  r  s   `` ``  @r;   _expand_inputs_for_generationz;Glm4vForConditionalGeneration._expand_inputs_for_generation  s     !l**w+	"Z
	" :,G !33KQ3GI2<@ 12: !pqq.I,WhJi.jL*+,&&r<   rf   )NNNNNNNNNNNr   )
NNNNTNNNNF)r&   FN)!rO   rP   rQ   _tied_weights_keysrL  r1   r   r4   r}  r  r   r   rL   r   r+  r0  r   r   rS   r   rM  r   rO  rJ   r`  rc  rx  rj   dictr   r   r  rT   rU   s   @r;   rQ  rQ    s"   *,VW  37\"..\ ((4/\ +,	\
 
+	+\ \  37U''U ((4/U +,	U
 
+	+U U ]G4 .2.204(,26*.,08<262648-.Y
##d*Y
 t+Y
 &&-	Y

 Y
 ((4/Y
   4'Y
 llT)Y
 #..5Y
 ((4/Y
 ((4/Y
 !??T1Y
 ell*Y
 +,Y
 
,	,Y
   5Y
|   $L$R .26*##d*6* ||d*6* 
u||U\\)	*	6*t #(-1	V'V' !V' ##d*	V' 
uc3h/	0V'r<   rQ  )rQ  r  r  r  r  )r
  )r&   )br  r  collections.abcr   dataclassesr   typingr   r   r4   torch.nnr2   torch.nn.functionalr   r   r    r	   r  activationsr
   cache_utilsr   r   
generationr   integrationsr   masking_utilsr   modeling_flash_attention_utilsr   modeling_layersr   modeling_outputsr   r   r   modeling_rope_utilsr   r   modeling_utilsr   r   processing_utilsr   utilsr   r   r   r   utils.deprecationr   utils.genericr   r    r!   r"   utils.output_capturingr#   vision_utilsr$   r%   configuration_glm4vr'   r(   r)   Moduler,   rW   rl   r|   r   r   r   rS   rL   r   r   r   rR   r   r   r  r+  r[  rd  rf  rv  r  r  r  r  r  r  rO  rQ  __all__r  r<   r;   <module>r     s  (   $ !        & ! . ) 7 / B 9 k k K F & a a 0  6 J P P Y'J299 J (J(fBII fBII (G GfRYY f"GBII GT(||+0<<>Cll
5<<%&	UU\\ 	U# 	U%,, 	U& %II%<<% 
% <<	%
 LL4'% % % '(%2P299 Pf1 >Jryy JZ6%PC) C)L)299 )$/6 /d 
06 0  0 2? 2 2(e
+ e
P e
) e
 e
P v
% v
 v
r 
0"8 0  0b'$8/ b'J xr<   