Ë
    T^(hÛ3  ã                   ó6  — d dl mZmZmZm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mZmZ dd	lmZ d
dlmZ  ej&                  e«      Z G d„ de«      Z G d„ dej.                  «      Z G d„ dej.                  «      Z G d„ de«      Z G d„ de«      ZddgZy)é    )ÚListÚOptionalÚTupleÚUnionN)Únné   )ÚACT2FN)Úis_torchdynamo_compilingÚloggingé   )ÚLlavaCausalLMOutputWithPastÚLlavaForConditionalGeneration)ÚMistralRMSNormé   )ÚMistral3Configc                   ó   — e Zd Zy)ÚMistral3RMSNormN©Ú__name__Ú
__module__Ú__qualname__© ó    úk/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/mistral3/modular_mistral3.pyr   r      ó   „ Ør   r   c                   óx   ‡ — e Zd ZdZdefˆ fd„Zdej                  dej                  dej                  fd„Zˆ xZ	S )ÚMistral3PatchMergerz<
    Learned merging of spatial_merge_size ** 2 patches
    Úconfigc                 ó"  •— t         ‰| �  «        || _        |j                  j                  }|j
                  | _        | j                  j                  j                  | _        t        j                  || j
                  dz  z  |d¬«      | _	        y )Nr   F©Úbias)
ÚsuperÚ__init__r   Úvision_configÚhidden_sizeÚspatial_merge_sizeÚ
patch_sizer   ÚLinearÚmerging_layer)Úselfr   r%   Ú	__class__s      €r   r#   zMistral3PatchMerger.__init__(   sr   ø€ Ü‰ÑÔØˆŒà×*Ñ*×6Ñ6ˆØ"(×";Ñ";ˆÔØŸ+™+×3Ñ3×>Ñ>ˆŒÜŸY™Y {°T×5LÑ5LÈaÑ5OÑ'OÐQ\ÐchÔiˆÕr   Úimage_featuresÚimage_sizesÚreturnc                 óÚ  — |D �cg c]&  }|d   | j                   z  |d   | j                   z  f‘Œ( }}|D ��cg c]
  \  }}||z  ‘Œ }}}|j                  d   }g }t        |j                  |«      «      D ]Á  \  }	}
||	   \  }}|
j	                  |||«      j                  ddd«      j                  d«      }t        j                  j                  j                  || j                  | j                  ¬«      }|j	                  || j                  dz  z  d«      j                  «       }|j                  |«       ŒÃ t        j                  |d¬«      }| j                  |«      }|S c c}w c c}}w )Nr   r   éÿÿÿÿr   )Úkernel_sizeÚstride©Údim)r'   ÚshapeÚ	enumerateÚsplitÚviewÚpermuteÚ	unsqueezeÚtorchr   Ú
functionalÚunfoldr&   ÚtÚappendÚcatr)   )r*   r,   r-   Ú
image_sizeÚhÚwÚtokens_per_imageÚdÚpermuted_tensorÚimage_indexÚimage_tokensÚ
image_gridÚgrids                r   ÚforwardzMistral3PatchMerger.forward1   sl  € àcnö
ØU_ˆZ˜‰]˜dŸo™oÑ-¨z¸!©}ÀÇÁÑ/OÒPð
ˆð 
ð /:×:¡d a¨˜A ›EÐ:ÐÑ:Ø× Ñ  Ñ$ˆàˆÜ)2°>×3GÑ3GÐHXÓ3YÓ)Zò 	)Ñ%ˆK˜à˜{Ñ+‰DˆAˆqØ%×*Ñ*¨1¨a°Ó3×;Ñ;¸A¸qÀ!ÓD×NÑNÈqÓQˆJÜ—8‘8×&Ñ&×-Ñ-Ø¨×(?Ñ(?È×H_ÑH_ð .ó ˆDð —9‘9˜Q ×!8Ñ!8¸!Ñ!;Ñ;¸RÓ@×BÑBÓDˆDØ×"Ñ" 4Õ(ð	)ô Ÿ™ ?¸Ô:ˆØ×+Ñ+¨NÓ;ˆØÐùò)
ùó ;s
   …+E"·E')
r   r   r   Ú__doc__r   r#   r;   ÚTensorrK   Ú__classcell__©r+   s   @r   r   r   #   s?   ø„ ñðj˜~õ jð e§l¡lð ÀÇÁð ÐRW×R^ÑR^÷ 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 )ÚMistral3MultiModalProjectorr   c                 ó^  •— t         ‰| �  «        t        |j                  j                  «      | _        t        |«      | _        t        |j                  t        «      rdnt        |j                  «      }t        j                  |j                  j                  |z  |j                  j                  |j                  ¬«      | _        t"        |j$                     | _        t        j                  |j                  j                  |j                  j                  |j                  ¬«      | _        y )Nr   r    )r"   r#   r   r$   r%   Únormr   Úpatch_mergerÚ
isinstanceÚvision_feature_layerÚintÚlenr   r(   Útext_configÚmultimodal_projector_biasÚlinear_1r	   Úprojector_hidden_actÚactÚlinear_2)r*   r   Únum_feature_layersr+   s      €r   r#   z$Mistral3MultiModalProjector.__init__J   sÞ   ø€ Ü‰ÑÔÜ# F×$8Ñ$8×$DÑ$DÓEˆŒ	Ü/°Ó7ˆÔä",¨V×-HÑ-HÌ#Ô"N™QÔTWÐX^×XsÑXsÓTtÐÜŸ	™	Ø× Ñ ×,Ñ,Ð/AÑAØ×Ñ×*Ñ*Ø×1Ñ1ô
ˆŒô
 ˜&×5Ñ5Ñ6ˆŒÜŸ	™	Ø×Ñ×*Ñ*¨F×,>Ñ,>×,JÑ,JÐQW×QqÑQqô
ˆ�r   r,   r-   c                 ó²   — | j                  |«      }| j                  ||«      }| j                  |«      }| j                  |«      }| j	                  |«      }|S )N)rS   rT   r[   r]   r^   )r*   r,   r-   Úhidden_statess       r   rK   z#Mistral3MultiModalProjector.forwardZ   sR   € ØŸ™ >Ó2ˆØ×*Ñ*¨>¸;ÓGˆØŸ™ nÓ5ˆØŸ™ Ó/ˆØŸ™ mÓ4ˆØÐr   )	r   r   r   r   r#   r;   rM   rK   rN   rO   s   @r   rQ   rQ   I   s*   ø„ ð
˜~õ 
ð  e§l¡lð ÀÇÁ÷ r   rQ   c                   ó   — e Zd Zy)ÚMistral3CausalLMOutputWithPastNr   r   r   r   rc   rc   c   r   r   rc   c            #       ó  — e Zd Zdej                  deeee   f   dej                  fd„Z		 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dde
ej                     de
ej                     de
ej                     de
ej                     d	e
eej                        d
e
ej                     de
eeee   f      de
ej                     de
e   de
e   de
e   de
e   de
ej                     deeej                  f   de
ej                     deeef   f d„Zy)Ú Mistral3ForConditionalGenerationÚpixel_valuesrV   r-   c                 ó|  — |j                  «       D ��ci c]  \  }}|€Œ	||“Œ }}} | j                  |f|ddœ|¤Ž}t        |t        «      r|j                  |   }n3|D �	cg c]  }	|j                  |	   ‘Œ }
}	t        j                  |
d¬«      }| j                  |j                  d«      |«      }|S c c}}w c c}	w )a=  
        Obtains image last hidden states from the vision tower and apply multimodal projection.

        Args:
            pixel_values (`torch.FloatTensor]` of shape `(batch_size, channels, height, width)`):
               The tensors corresponding to the input images.
            vision_feature_layer (`Union[int, List[int]]`):
                The index of the layer to select the vision feature. If multiple indices are provided,
                the vision feature of the corresponding indices will be concatenated to form the
                vision features.
            image_sizes (`torch.Tensor`):
                Tensor containing the image sizes as returned by the processor.
        Returns:
            image_features (`torch.Tensor`): Image feature tensor of shape `(num_images, image_length, embed_dim)`).
        T)r-   Úoutput_hidden_statesr0   r3   r   )	ÚitemsÚvision_towerrU   rW   ra   r;   r@   Úmulti_modal_projectorÚsqueeze)r*   rf   rV   r-   ÚkwargsÚkÚvÚimage_outputsÚselected_image_featureÚ	layer_idxÚhs_poolr,   s               r   Úget_image_featuresz3Mistral3ForConditionalGeneration.get_image_featuresh   sÉ   € ð, $*§<¡<£>×C™4˜1˜a°Q±]�!�Q‘$ÐCˆÑCà)˜×)Ñ)¨,ÐuÀKÐfjÑuÐntÑuˆô Ð*¬CÔ0Ø%2×%@Ñ%@ÐAUÑ%VÑ"àOcÖdÀ)�}×2Ñ2°9Ó=ÐdˆGÐdÜ%*§Y¡Y¨w¸BÔ%?Ð"à×3Ñ3Ð4J×4RÑ4RÐSTÓ4UÐWbÓcˆØÐùó Dùò es   ”
B3ŸB3Á!B9NÚ	input_idsÚattention_maskÚposition_idsÚpast_key_valuesÚinputs_embedsÚlabelsÚ	use_cacheÚoutput_attentionsrh   Úreturn_dictÚcache_positionÚlogits_to_keepr.   c                 óô  — |
�|
n| j                   j                  }
|�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|du |duz  rt        d«      ‚|�|�t        d«      ‚|€ | j                  «       |«      }|��#| j                  |||¬«      }|| j                   j                  k(  j                  d«      }|j                  |«      j                  |j                  «      }t        «       s{||   j                  «       |j                  «       k7  rW|| j                   j                  k(  j                  «       }|j                   d   |j                   d   z  }t        d|› d	|› �«      ‚|j                  |j                  |j"                  «      }|j%                  ||«      } | j&                  d|||||	|
||||d
œ
|¤Ž}|d   }d}|��<|�¥|dd…|j                   d   dz
   d…f   j                  |j                  «      }|ddd…dd…f   |j                  |j                  «      dk7     j)                  «       }|ddd…f   |j                  |j                  «      dk7     j)                  «       }n1|ddd…dd…f   j)                  «       }|ddd…f   j)                  «       }t+        j,                  «       } ||j/                  d|j1                  d«      «      |j/                  d«      j                  |j                  «      «      }|s|f|dd z   }|�|f|z   S |S t3        |||j4                  |j6                  |j8                  |�¬«      S d¬«      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]`.

            logits_to_keep (`int` or `torch.Tensor`, *optional*):
                If an `int`, compute logits for the last `logits_to_keep` tokens. If `0`, calculate logits for all
                `input_ids` (special case). Only last token logits are needed for generation, and calculating them only for that
                token can save memory, which becomes pretty significant for long sequences or large vocabulary size.
                If a `torch.Tensor`, must be 1D corresponding to the indices to keep in the sequence length dimension.
                This is useful when using packed tensor format (single dimension for batch and sequence length).


        Returns:

        Example:

        ```python
        >>> from PIL import Image
        >>> import requests
        >>> from transformers import AutoProcessor, Mistral3ForConditionalGeneration

        >>> model = Mistral3ForConditionalGeneration.from_pretrained("mistralai/Mistral-Small-3.1-24B-Instruct-2503")
        >>> processor = AutoProcessor.from_pretrained("mistralai/Mistral-Small-3.1-24B-Instruct-2503")

        >>> prompt = "<s>[INST][IMG]What is the image?[/INST]"
        >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
        >>> image = Image.open(requests.get(url, stream=True).raw)

        >>> inputs = processor(images=image, text=prompt, return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(**inputs, max_new_tokens=15)
        >>> processor.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "What is the image?The image depicts two cats lying on a pink blanket."
        ```Nz:You must specify exactly one of input_ids or inputs_embedszdYou cannot specify both pixel_values and inputs_embeds at the same time, and must specify either one)rf   rV   r-   r0   r   r   z6Image features and image tokens do not match: tokens: z, features )
rv   rw   rx   ry   r{   r|   rh   r}   r~   r   .)ÚlossÚlogitsrx   ra   Ú
attentionsÚimage_hidden_statesr   )r   r|   rh   Úuse_return_dictrV   Ú
ValueErrorÚget_input_embeddingsrt   Úimage_token_indexr:   Ú	expand_asÚtoÚdevicer
   ÚnumelÚsumr5   ÚdtypeÚmasked_scatterÚlanguage_modelÚ
contiguousr   ÚCrossEntropyLossr8   Úsizerc   rx   ra   rƒ   )r*   ru   rf   rv   rw   rx   ry   rV   rz   r{   r|   rh   r}   r~   r   r-   Ú	lm_kwargsr,   Úspecial_image_maskÚn_image_tokensÚn_image_featuresÚoutputsr‚   r�   Úshift_attention_maskÚshift_logitsÚshift_labelsÚloss_fctÚoutputs                                r   rK   z(Mistral3ForConditionalGeneration.forwardŒ   sâ  € ðr 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð ˜Ð -°tÐ";Ò<ÜÐYÓZÐZàÐ#¨Ð(AÜØvóð ð Ð Ø7˜D×5Ñ5Ó7¸	ÓBˆMàÑ#Ø!×4Ñ4Ø)Ø%9Ø'ð 5ó ˆNð #,¨t¯{©{×/LÑ/LÑ"L×!WÑ!WÐXZÓ![ÐØ!3×!=Ñ!=¸mÓ!L×!OÑ!OÐP]×PdÑPdÓ!eÐÜ+Ô-°-Ð@RÑ2S×2YÑ2YÓ2[Ð_m×_sÑ_sÓ_uÒ2uØ"+¨t¯{©{×/LÑ/LÑ"L×!QÑ!QÓ!S�Ø#1×#7Ñ#7¸Ñ#:¸^×=QÑ=QÐRSÑ=TÑ#TÐ Ü ØLÈ^ÐL\Ð\gÐhxÐgyÐzóð ð ,×.Ñ.¨}×/CÑ/CÀ]×EXÑEXÓYˆNØ)×8Ñ8Ð9KÈ^Ó\ˆMà%�$×%Ñ%ð 
Ø)Ø%Ø+Ø'ØØ/Ø!5Ø#Ø)Ø)ñ
ð ñ
ˆð ˜‘ˆàˆØÑàÐ)ð (6²a¸6¿<¹<È¹?ÈQÑ;NÐ9OÑ9QÐ6QÑ'R×'UÑ'UÐV\×VcÑVcÓ'dÐ$Ø% c¨3¨B¨3² kÑ2Ð3G×3JÑ3JÈ6Ï=É=Ó3YÐ]^Ñ3^Ñ_×jÑjÓl�Ø% c¨1©2 g™Ð/C×/FÑ/FÀvÇ}Á}Ó/UÐYZÑ/ZÑ[×fÑfÓh‘à% c¨3¨B¨3² kÑ2×=Ñ=Ó?�Ø% c¨1©2 g™×9Ñ9Ó;�ä×*Ñ*Ó,ˆHÙØ×!Ñ! " l×&7Ñ&7¸Ó&;Ó<¸l×>OÑ>OÐPRÓ>S×>VÑ>VÐWc×WjÑWjÓ>kóˆDñ Ø�Y ¨¨ Ñ,ˆFØ'+Ð'7�D�7˜VÑ#ÐC¸VÐCä-ØØØ#×3Ñ3Ø!×/Ñ/Ø×)Ñ)Ø2>Ð2J ô
ð 	
ð QUô
ð 	
r   )NNNNNNNNNNNNNr   N)r   r   r   r;   ÚFloatTensorr   rW   r   rM   rt   r   Ú
LongTensorÚboolr   rc   rK   r   r   r   re   re   g   s½  „ ð"à×'Ñ'ð"ð $ C¨¨c© NÑ3ð"ð —\‘\ó	"ðL 15Ø48Ø15Ø37Ø=AØ59Ø@DØ-1Ø$(Ø,0Ø/3Ø&*Ø59Ø34Ø.2ñ!L
à˜E×,Ñ,Ñ-ðL
ð ˜u×0Ñ0Ñ1ðL
ð ! §¡Ñ.ð	L
ð
 ˜u×/Ñ/Ñ0ðL
ð " $ u×'8Ñ'8Ñ"9Ñ:ðL
ð   × 1Ñ 1Ñ2ðL
ð ' u¨S°$°s±)¨^Ñ'<Ñ=ðL
ð ˜×)Ñ)Ñ*ðL
ð ˜D‘>ðL
ð $ D™>ðL
ð ' t™nðL
ð ˜d‘^ðL
ð ! ×!1Ñ!1Ñ2ðL
ð ˜c 5§<¡<Ð/Ñ0ðL
ð  ˜eŸl™lÑ+ð!L
ð$ 
ˆuÐ4Ð4Ñ	5ô%L
r   re   ÚMistral3PreTrainedModel)Útypingr   r   r   r   r;   r   Úactivationsr	   Úutilsr
   r   Úllava.modeling_llavar   r   Úmistral.modeling_mistralr   Úconfiguration_mistral3r   Ú
get_loggerr   Úloggerr   ÚModuler   rQ   rc   re   Ú__all__r   r   r   ú<module>r¬      s�   ð÷  0Ó /ã Ý å !ß 6ß ]Ý 5Ý 2ð 
ˆ×	Ñ	˜HÓ	%€ô	�nô 	ô#˜"Ÿ)™)ô #ôL "§)¡)ô ô4	Ð%@ô 	ôq
Ð'Dô q
ðj Ø&ð�r   