a
    d                  V   @   s  U d dl Z d dlZd dlZd dlZd dlZd dlZd dlZd dlZd dlZd dl	m
Z
mZmZmZmZmZmZ d dlZd dlmZ d dlmZmZmZ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! dd
l"m#Z#m$Z$m%Z%m&Z&m'Z'm(Z(m)Z)m*Z*m+Z+m,Z,m-Z-m.Z.m/Z/m0Z0m1Z1m2Z2m3Z3m4Z4m5Z5 ddl6m7Z7m8Z8m9Z9m:Z:m;Z; e: rVd dl<m=Z= e>e?Z@ejABddC e7v ZDdee eeeEeeE f  eeE dddZFg dZGg ZHeGD ]6ZIeJeIeKreHLeFf i eI neHLeFeI qg dZMeNeOePeHeM ZQdd ZRdddZSdd ZTdd ZUdd ZVd d! ZWd"d# ZXdd$d%ZYd&d' ZZdd(d)d*Z[d+d, Z\d-d. Z]ddd(d/d0Z^ddd(d1d2Z_d3dd4d5d6Z`dd(d7d8Zad9d: Zbdd(d;d<Zcdd(d=d>Zdd3d3dd?d@dAZed3d3dd?dBdCZfdDdE ZgdFdG ZhdddHdIdJZidd(dKdLZjdMdN ZkdddOdPdQZldRdS ZmddTdUZndVdW ZodXdY ZpdZd[ Zqd\d] Zrdd^d_Zsdd`daZtdbdc Zuddde Zvdfdg ZwddidjZxdkdl Zydmdn Zzdodp Z{dqdr Z|ejj}eRejj~jeSejjeTejjeUejjeVejeWejj~jeYejjeXejeZeje[eje\eje]eje^eje_eje`ejeaejjebejecejedejeeejjefejegejjehejeiejenejeoejjepejejejjekejelejjemejjeqejjerejesejjetejeuejjevejewejj~jexejjeyejjezejje{eje|i+Zeeef eds< G dtdu dueZG dvdw dweZG dxdy dyeZdzd{ Zd|d} Zdeeeee  dddZG dd deZejeeE dddZedddZddefeeeeE  eee edddZdS )    N)AnyCallableDictListOptionalTypeUnion)nn)GraphGraphModuleProxyTracer)compatibilityParameterProxy   )PretrainedConfigPreTrainedModellogging)
get_values),MODEL_FOR_AUDIO_CLASSIFICATION_MAPPING_NAMES MODEL_FOR_BACKBONE_MAPPING_NAMES!MODEL_FOR_CAUSAL_LM_MAPPING_NAMESMODEL_FOR_CTC_MAPPING_NAMES3MODEL_FOR_DOCUMENT_QUESTION_ANSWERING_MAPPING_NAMES,MODEL_FOR_IMAGE_CLASSIFICATION_MAPPING_NAMES-MODEL_FOR_MASKED_IMAGE_MODELING_MAPPING_NAMES!MODEL_FOR_MASKED_LM_MAPPING_NAMES'MODEL_FOR_MULTIPLE_CHOICE_MAPPING_NAMES0MODEL_FOR_NEXT_SENTENCE_PREDICTION_MAPPING_NAMES#MODEL_FOR_PRETRAINING_MAPPING_NAMES*MODEL_FOR_QUESTION_ANSWERING_MAPPING_NAMES-MODEL_FOR_SEMANTIC_SEGMENTATION_MAPPING_NAMES,MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING_NAMES/MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING_NAMES(MODEL_FOR_SPEECH_SEQ_2_SEQ_MAPPING_NAMES,MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING_NAMES6MODEL_FOR_ZERO_SHOT_IMAGE_CLASSIFICATION_MAPPING_NAMESMODEL_MAPPING_NAMES)ENV_VARS_TRUE_VALUESTORCH_FX_REQUIRED_VERSIONget_torch_versionis_peft_availableis_torch_fx_available)	PeftModelZFX_DEBUG_MODE )
model_namesupported_tasksreturnc                 C   s|   t ttttttttt	t
ttttttttd}|d u r<| }t|trL|g}g }|D ]"}|| | d }|rT|| qT|S )N)defaultZpretrainingznext-sentence-predictionz	masked-lmz	causal-lmz
seq2seq-lmzspeech-seq2seqzmultiple-choicezdocument-question-answeringzquestion-answeringzsequence-classificationztoken-classificationzmasked-image-modelingzimage-classificationzzero-shot-image-classificationZctczaudio-classificationzsemantic-segmentationZbackbone)r(   r    r   r   r   r#   r%   r   r   r!   r$   r&   r   r   r'   r   r   r"   r   keys
isinstancestrgetappend)r0   r1   Ztask_mappingZmodel_class_namesZtask
class_name r:   ^/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/transformers/utils/fx.py%_generate_supported_model_class_namesI   s<    
r<   ))ZaltclipZalbertZbartZbertZ
blenderbotzblenderbot-smallZbloomZclipZconvnextZdebertaz
deberta-v2Z
distilbertz
donut-swinZelectraZgpt2Zgpt_neoZgptjZhubertZlayoutlmZlxmertZm2m_100ZmarianZmbartzmegatron-bertZ
mobilebertZmt5ZnezhaoptZpegasusZplbartZresnetZrobertaZ	segformerZspeech_to_textZspeech_to_text_2ZswinZt5ZtrocrZvitZxglmZwav2vec2)ZCLIPTextModelZCLIPTextModelWithProjectionZCLIPVisionModelZCLIPVisionModelWithProjectionZAltCLIPTextModelZAltCLIPVisionModelZGitVisionModelGPT2DoubleHeadsModelZSpeech2Text2DecoderZTrOCRDecoderPeftModelForCausalLMPeftModelForSeq2SeqLMc                 C   s*   t jg |j| jjd R d| jjdS Nmeta)devicedtype)torchemptyshapeweightrE   selfinputr:   r:   r;   torch_nn_embedding   s    rM          @Fc                 C   s&   t jg | j|jd R d|jdS rA   )rF   rG   rH   rE   )rL   rI   Zpadding_idxZmax_normZ	norm_typeZscale_grad_by_freqsparser:   r:   r;   torch_nn_functional_embedding   s    rP   c                 C   s   |S Nr:   rJ   r:   r:   r;   torch_nn_layernorm   s    rR   c                 C   s   |S rQ   r:   rJ   r:   r:   r;   torch_nn_groupnorm   s    rS   c                 C   s    t j|jd d | jf ddS )NrB   rC   rD   )rF   rG   rH   Zout_featuresrJ   r:   r:   r;   torch_nn_linear   s    rU   c                 C   s   | S rQ   r:   xr:   r:   r;   
torch_relu   s    rX   c                 C   s   |S rQ   r:   )rK   rW   r:   r:   r;   torch_nn_relu   s    rY   c                 C   s   |st d| S )Nz>Don't support in-place functional.relu for MetaTensor analysis
ValueError)rW   Zinplacer:   r:   r;   torch_nn_functional_relu   s    r\   c                 C   s$   | j dd|j dd |j dd S NrC   rT   to)	conditionrW   yr:   r:   r;   torch_where   s    rb   outc                C   s   |d urt d| S )Nz2Don't support in-place abs for MetaTensor analysisrZ   )rL   rd   r:   r:   r;   	torch_abs   s    re   c                  O   s   t | }d}|dkr"d}| d }n|dkr4| \}}n
| \}}}t|trPt|}t|trbt|}t|trtt|}|d|}|d}tj|| | |ddS )N   r   r   steprE   rC   rE   rD   )lenr5   floatintr7   rF   rG   )argskwargsnrg   startendrE   r:   r:   r;   torch_arange   s"    






rq   c                  O   sX   t | } t| d tjr4| d jtdkr4d| d< t|}|dd  tj| i |S )Nrf   rC   rD   )listr5   rF   TensorrD   dictpopfull)rl   rm   Zkwargs_without_devicer:   r:   r;   
torch_full   s    $rw   c                   s    d u r|d u rd  d u r(|d ur(|  dk r@| d      dd | D }t|d }t fdd|D }|d   |g | d d   }tj|ddS )	Nr   c                 S   s   g | ]
}|j qS r:   )rH   ).0tr:   r:   r;   
<listcomp>      ztorch_cat.<locals>.<listcomp>c                 3   s   | ]}|  V  qd S rQ   r:   )rx   rH   dimr:   r;   	<genexpr>  r{   ztorch_cat.<locals>.<genexpr>rf   rC   rT   )r}   rr   sumrF   rG   )tensorsr}   axisrd   ZshapesrH   Zconcatenated_dimZfinal_shaper:   r|   r;   	torch_cat  s    "r   c                C   sp   |d u r|d u rd}|d u r(|d ur(|}|dk rD| d   d | }t| d j}||t|  tj|ddS Nr   rf   rC   rT   )r}   rr   rH   insertri   rF   rG   )r   r}   r   rd   rH   r:   r:   r;   torch_stack  s    r   rf   )alphard   c          	      C   s   t | tjstj|ddS t |tjs4tj| ddS t|  | }t| jdg||     }t|jdg||    }g }t|D ]}|	t|| ||  qtj
|ddS )NrC   rT   rf   )r5   rF   rs   
empty_likemaxr}   rr   rH   ranger8   rG   )	rL   otherr   rd   
max_lengthZinput_shapeZother_shaperH   ir:   r:   r;   	torch_add   s    r   c                C   s   t | ||dS )Nrc   )r   )rL   r   rd   r:   r:   r;   	torch_mul.  s    r   c                 C   s
   t | |S rQ   )r   )rK   r   r:   r:   r;   torch_tensor_mul2  s    r   c          
      C   s  |   }|  }d }|dkr,|dkr,d }nL|dkrT|dkrT| d|df}n$|dkrt|dkrt|df}n|dkr|dkr| df}nt|   |  }t| j}t|j}|dkrdg| }|dkr|d dg||  t| j }dg||  t|j }g }t|D ]}	|t||	 ||	  q|d |d< |d |d< |dkrd|d |dkrx|d |d u rtj	dddS tj
|d	diS )
Nrf   r   r   rB   g        rC   rT   rD   )r}   sizer   rr   rH   r8   r   ru   rF   tensorrG   )
rL   r   rd   d1Zd2rH   r   Zshape1Zshape2r   r:   r:   r;   torch_matmul6  s@    








r   c                C   s:   |d urt d| j\}}}|j\}}}tj|||ddS )Nz2Don't support in-place bmm for MetaTensor analysisrC   rT   )r[   rH   rF   rG   )rL   Zmat2rd   
batch_sizern   m_pr:   r:   r;   	torch_bmmZ  s
    r   betar   rd   c                C   s   |d urt dt||S )Nz6Don't support in-place baddbmm for MetaTensor analysis)r[   r   )rL   batch1batch2r   r   rd   r:   r:   r;   torch_baddbmmb  s    r   c                C   s   t | |||||dS )Nr   )r   )rK   r   r   r   r   rd   r:   r:   r;   torch_tensor_baddbmmh  s    r   c                 G   s&   dd |D }t j| g|R  dS )Nc                 s   s   | ]}t j|d dV  qdS )cpurT   N)rF   r   )rx   operandr:   r:   r;   r~   n  r{   ztorch_einsum.<locals>.<genexpr>rC   )rF   einsumr_   )ZequationZoperandsZconcrete_operandsr:   r:   r;   torch_einsuml  s    r   c                 G   s:   t | j}t|D ]\}}||  |9  < qtj|ddS r]   )rr   rH   	enumeraterF   rG   )rK   sizesrH   r   rW   r:   r:   r;   torch_tensor_repeatr  s    
r   )r}   output_sizec                 G   s   t |}|dkr,|d ur|n
|d  g}nt|d j}| d u rb|dkrT|d } nt|g}d} |d }t|tst|dkr||   t|9  < n|d ur|n| || < tj|ddiS )Nrf   r   r   rD   rC   )	ri   r   rr   rH   r5   rk   rF   ZnumelrG   )r}   r   rl   num_argsrH   Zrepeatsr:   r:   r;   torch_repeat_interleavey  s    

r   c                C   s&   t | j}t|||< tj|ddiS NrD   rC   )rr   rH   ri   rF   rG   )rL   r}   indexrd   rH   r:   r:   r;   torch_index_select  s    
r   c                 C   s   t | ||S rQ   )r   rK   r}   r   r:   r:   r;   torch_tensor_index_select  s    r   )sparse_gradrd   c                C   s(   t | j}|j| ||< tj|ddiS r   )rr   rH   rF   rG   )rL   r}   r   r   rd   rH   r:   r:   r;   torch_gather  s    
r   c                 C   s   t | ||S rQ   )r   r   r:   r:   r;   torch_tensor_gather  s    r   c                 C   s   | S rQ   r:   )rL   Zshiftsdimsr:   r:   r;   
torch_roll  s    r   c                 C   s   | S rQ   r:   )rL   r   r:   r:   r;   
torch_flip  s    r   c                 C   s   | S rQ   r:   )rK   r   r:   r:   r;   torch_tensor_flip  s    r   c                 C   s   |j d }d }| j}|dkr d}|dkr2t|j }|d u rt|j }t|d|d   | jd | jd d   d | jd  d }||d< | j|d< t	j
|d	d
S )NrB   validr   r   samer   r   rf   r   rC   rT   rH   paddingrr   mathfloorZdilationZkernel_sizeZstrideZout_channelsrF   rG   )rK   rL   Zl_inrH   r   Zl_outr:   r:   r;   torch_nn_conv1d  s    


8
r   c                 C   s   |j dd  \}}d }| j}|dkr(d}|dkr:t|j }|d u rt|j }t|d|d   | jd | jd d   d | jd  d }t|d|d   | jd | jd d   d | jd  d }||g|dd < | j|d< t	j
|d	d
S )Nr   r   r   r   r   r   rf   rC   rT   r   )rK   rL   Zh_inZw_inrH   r   Zh_outZw_outr:   r:   r;   torch_nn_conv2d  s$    

88
r   c                 C   sr   t | j}|d ur>|dk r&|  | }|| dkrd|| n&g }|D ]}|dkrTqF|| qF|}tj|ddS r   )rr   rH   r}   ru   r8   rF   rG   )rL   r}   rH   Z	new_shapeZ	dim_valuer:   r:   r;   torch_squeeze  s    
r   c                 C   s
   t | |S rQ   )r   rK   r}   r:   r:   r;   torch_tensor_squeeze  s    r   c                 C   s<   t | j}|dk r"|  d | }||d tj|ddS r   )rr   rH   r}   r   rF   rG   )rL   r}   rH   r:   r:   r;   torch_unsqueeze  s
    
r   c                 C   s
   t | |S rQ   )r   r   r:   r:   r;   torch_tensor_unsqueeze  s    r   c                 K   sH   t jt j| ddfi |}t|t jr2|dS tt|dd S d S )Nr   rT   rC   c                 S   s
   |  dS )NrC   r^   rV   r:   r:   r;   <lambda>  r{   z*torch_unique_consecutive.<locals>.<lambda>)rF   unique_consecutiveZ
zeros_liker5   rs   r_   tuplemap)rL   rm   outputr:   r:   r;   torch_unique_consecutive  s    
r   rB   c                 C   s.   |dk rt dt| j|g }tj|ddS )Nr   zEDon't support automatic num_classes inference for MetaTensor analysisrC   rT   )r[   rr   rH   rF   rG   )r   Znum_classesrH   r:   r:   r;   torch_nn_functional_one_hot  s    r   c                 C   s$   | j dkr|j}nd}tj|ddS Nnone)rf   rC   rT   Z	reductionrH   rF   rG   rK   rL   targetrH   r:   r:   r;   torch_nn_mseloss  s    
r   c                 C   s$   | j dkr|j}nd}tj|ddS r   r   r   r:   r:   r;   torch_nn_crossentropyloss  s    
r   c                 C   s$   | j dkr|j}nd}tj|ddS r   r   r   r:   r:   r;   torch_nn_bcewithlogitsloss  s    
r   c                 C   s^   dd }t | tjrRt |tr.tt||}n||}ttj| dd|dS t| |S )Nc                 S   sH   t | tjrDtj| dd}|jtjtjtjtjfv r@|	tj
}|S | S )Nr   rT   )r5   rF   rs   Z	ones_likerE   float16float32float64int32r_   int64)ry   Zconcreter:   r:   r;   to_concrete  s    z%operator_getitem.<locals>.to_concreter   rT   rC   )	r5   rF   rs   r   r   operatorgetitemr   r_   )abr   r:   r:   r;   operator_getitem  s    
r   _MANUAL_META_OVERRIDESc                       sh   e Zd ZdZdd Zedd Zedd Z fdd	Z fd
dZ	dd Z
dd Z fddZ  ZS )HFProxyzI
    Proxy that uses metadata to handle data-dependent control-flow.
    c                 C   s
   || _ d S rQ   )	_metadata)rK   metadatar:   r:   r;   install_metadatac  s    zHFProxy.install_metadatac                 C   s   | j dd| fi S )Ncall_methodr   )tracercreate_proxyrK   r:   r:   r;   rH   f  s    zHFProxy.shapec                 C   s
   t | dS )NrD   )MetaDeviceAttributer   r:   r:   r;   rD   j  s    zHFProxy.devicec                    s(   t | dr| jd urt| jS t  S Nr   )hasattrr   ri   super__len__r   	__class__r:   r;   r   p  s    
zHFProxy.__len__c                    s$   t | dr| jd ur| jS t  S r   )r   r   r   __bool__r   r   r:   r;   r   u  s    zHFProxy.__bool__c                 C   s   |dkr|  |S t| |S r   )__getattribute__HFAttribute)rK   kr:   r:   r;   __getattr__z  s    
zHFProxy.__getattr__c                 C   s   | j dtj| ||fi S Ncall_function)r   r   r   setitem)rK   indicesvaluesr:   r:   r;   __setitem__  s    zHFProxy.__setitem__c                    s*   t | dr| jd ur|| jv S t |S r   )r   r   r   __contains__)rK   keyr   r:   r;   r     s    
zHFProxy.__contains__)__name__
__module____qualname____doc__r   propertyrH   rD   r   r   r   r   r   __classcell__r:   r:   r   r;   r   ^  s   

r   c                   @   s.   e Zd ZedddZedd Zdd ZdS )	r   )attrc                 C   s>   || _ || _|j| _d | _t| j dr:| t| j j| d S r   )rootr  r   _noder   r   getattrr   )rK   r  r  r:   r:   r;   __init__  s    zHFAttribute.__init__c                 C   s0   | j d u r*| jdtj| j| jfi j| _ | j S r   )r  r   r   builtinsr  r  r  noder   r:   r:   r;   r    s    
 zHFAttribute.nodec                 O   s   | j d| j| jf| |S )Nr   )r   r   r  r  )rK   rl   rm   r:   r:   r;   __call__  s    zHFAttribute.__call__N)r   r   r   r6   r  r   r  r	  r:   r:   r:   r;   r     s   	
r   c                   @   s   e Zd ZdS )r   N)r   r   r   r:   r:   r:   r;   r     s   r   c                 C   sH   t | trdS t | tjjrDt | tr0t| ds>td|  | jS | S )z\Returns the underlying metadata for HFProxies, and behaves like the identity for the others.rC   r   zNo metadata was found for )	r5   r   rF   fxr   r   r   RuntimeErrorr   vr:   r:   r;   _proxies_to_metas  s    
r  c                    s   t   fdd}| fS )Nc                     s\   d   fdd}t jj| | t jj||  d urJ jd| |S | i |S d S )Nc                    s   t | tr|  d S rQ   r5   r   r  proxyr:   r;   check_has_proxy  s    
zB_gen_constructor_wrapper.<locals>.wrapper.<locals>.check_has_proxyr   )rF   r
  r  map_aggregater   r   )rl   rm   r  r   r  r;   wrapper  s    z)_gen_constructor_wrapper.<locals>.wrapper)	functoolswraps)r   r  r:   r  r;   _gen_constructor_wrapper  s    r  
      )lowhighforbidden_valuesc                 C   s2   |d u rg }t | |}||v r.t | |}q|S rQ   )randomrandint)r  r  r  valuer:   r:   r;   _generate_random_int  s    r!  c                       sz  e Zd ZU dZdZeed< dZeed< g dZe	 s:e
fne
efZefdf fdd	Ze
eee eeejf d	d
dZd+ fdd	Zdd Zeeeeef dddZ fddZdd Zd,eejjedef f e eeef  e eeef  ee!d fddZ"ejedddZ#ejeddd Z$ejed fd!d"Z%ejjeed# fd$d%Z&e'dd&d'ed(d)d*Z(  Z)S )-HFTracerz
    Tracer that is able to symbolically trace models from the library. To do that, it uses the HFProxy instead of the
    regular PyTorch torch.fx.Proxy.
    Tproxy_buffer_attributesallow_insert_stateless_mods)
arangezerosZonesrv   Z	full_likeZeyerG   r   clampZfinfor:   c                    s2   t  j||d t s.tdt  dt dd S )N)autowrap_modulesautowrap_functionsz6Found an incompatible version of torch. Found version z, but only version z is supported.)r   r  r-   ImportErrorr+   r*   )rK   r(  r)  r   r:   r;   r    s    
zHFTracer.__init__)model
input_namerH   r2   c                 C   sB  t |d|jj}|j}i }|dv r|d }|g tttttttttt	v rvt
j|t
j|d|d< q>|g ttttdv rt
j|t
j|d|d< t
j|t
j|d|d< q>|ttv r|t|jd	r|jjd
u rtd|jjdkr||jjf}t
j}	nR|jjdkr.|f}t
j}	n6|jjdkrP||jjf}t
j}	ntd|jj dt
j||	|d|d< n|g ttttttttttttdddv rt
j|t
j|d|d< n@|g ttv rt
j|t
j|d|d< ntd| d| dn d|v r|d }t |jdd
}
|
d
u rt|jdrb|jjj}
n&t|jdr||jj j}
nt! t! f}
t |jdd}t"|
t#j$j%s|
|
f}
|
\}}t
j||||t
j|d||< nhd|v rt
jg |dR t
j&|d||< n8d|v r:t
jg ||jj'R t
j&|d||< nd |v rft
j||jj(g t
j&|d||< nd!|v rt
j||jj)g t
j&|d||< nd"|v rt
j|t
j&|d||< nd#|v r|\}}t!d$d%d&}t
j||t
j&|d||< nPd'|v sd(|v rt
j|t
j|d||< n$||jj*g }t
j|t
j&|d||< |S ))z4Generates dummy input for model inference recording.class_for_deserialization)labelsstart_positionsend_positionsr   rh   r.  ZXLNetForQuestionAnsweringr/  r0  problem_typeNzCould not retrieve the problem type for the sequence classification task, please set model.config.problem_type to one of the following values: "regression", "single_label_classification", or "multi_label_classification".Z
regressionZsingle_label_classificationZmulti_label_classificationzExpected model.config.problem_type to be either: "regression", "single_label_classification", or "multi_label_classification", but "z" was provided.r>   r?   r@   z!Generating the dummy input named z for z is not supported yet.Zpixel_values
image_sizevision_configencodernum_channels   Zbbox   Zinput_featuresZvisual_featsZ
visual_posinputsZinput_valuesi'  i N  r  r  maskids)+r  r   r   rD   r   r   r   r   r   r   rF   r&  longr!   r   r$   r   configr1  r[   Z
num_labelsr   r    r&   r   r   r#   r"   r   NotImplementedErrorr3  r2  r4  r!  r5   collectionsabcIterablerj   Zinput_feat_per_channelZvisual_feat_dimZvisual_pos_dimZhidden_size)rK   r+  r,  rH   Zmodel_class_namerD   Zinputs_dictr   Zlabels_shapeZlabels_dtyper2  r5  heightwidthr   Z
seq_lengthZshape_with_hidden_sizer:   r:   r;   _generate_dummy_input  s    
	


&




zHFTracer._generate_dummy_inputNc                    sX  t  |||||||}|dkr>|| jv r>|| j|  |S || jv rXd|v rXd|d< ztjj|t	}	tjj|t	}
|dkrt
||}||	i |
}t|tjr|jdd}n0|dkrt|	d j|}t
||}||	i |
}n|dkrxt| d	st|  d
d| _zT| j|}t|}|t
v rTt
| |g|	R i |
}n| j|	i |
}W d| _nd| _0 nr|dkrd| _zP| j}|d}|D ]}t||}qt|tjr|jdd}n|}W d| _nd| _0 n|W S t|tstd|| W nH tyR } z.tr>td| d| d|  W Y d }~n
d }~0 0 |S )NplaceholderrD   rC   r   rT   r   r   call_moduleorig_forwardz/ does not have an attribute called orig_forwardTFget_attr.z"Don't support composite output yetzCould not compute metadata for z target z: )r   r   	meta_argsr   orig_fnsrF   r
  r  r  r  r   r7   r5   rs   r_   r  r   r   AttributeError_disable_module_getattrr  Zget_submoduletyperG  splitr   r[   	Exception_IS_IN_DEBUG_MODEwarningswarn)rK   kindr   rl   rm   nameZ	type_exprproxy_factory_fnrvZ
args_metasZkwargs_metasZmeta_targetZmeta_outmethodmodmod_typeZattr_itrZatomsZatomer   r:   r;   r   l  sb    




2zHFTracer.create_proxyc                    s   t  ddr|S  fdd}t|tjjrH|| j |}|d urH|S  jrxt|tjrx|| j	 |}|d urx|S |S d S )NrM  Fc                    s   |D ]x\} |u r|vrpi }dt jjv rPjs<d n fdd|d< jddi fi |}||< |   S qd S )NrV  c                    s   t |  S rQ   r   )r  )attr_valrn   rK   r:   r;   r     r{   zLHFTracer._module_getattr.<locals>.maybe_get_proxy_for_attr.<locals>.<lambda>rH  r:   )inspect	signaturer   
parametersZparam_shapes_constant)r\  Zcollection_to_searchparameter_proxy_cacher   rm   Z	val_proxyr   )r\  rn   r;   maybe_get_proxy_for_attr  s    z:HFTracer._module_getattr.<locals>.maybe_get_proxy_for_attr)
r  r5   rF   r	   	Parameterr  Znamed_parametersr#  rs   Znamed_buffers)rK   r  r\  r`  ra  Zmaybe_parameter_proxyZmaybe_buffer_proxyr:   r   r;   _module_getattr  s     zHFTracer._module_getattr)r  r\  r`  c                 C   s   |  |||S rQ   )rc  )rK   r  r\  r`  r:   r:   r;   r    s    zHFTracer.getattrc                    s   || _ t ||||S rQ   )rG  r   rF  )rK   r   forwardrl   rm   r   r:   r;   rF    s    zHFTracer.call_modulec                 C   s
   t || S rQ   )r   )rK   r  r:   r:   r;   r    s    zHFTracer.proxy.)r  concrete_argsdummy_inputs6complete_concrete_args_with_inputs_not_in_dummy_inputsr2   c                    s  t t|tjjr|jn|} du r*i  dur|r|j D ]0}|j	v rPq@|j
t jju r@td|j	 dq@  fdd|j D  |j    }t }t }	||	g}
|jjttv rtddd}|
d	| durtni }|D ]V}||v rqt|| js(t|jd
r>|| |||
 qtd| dqdd | D }|j D ]2}|jt jjkrl|j	|vrli |d|j	 < ql|| _ dd | j!D | _"t# | _$| j" D ]&\}\}}t%t|| | j$&| qz<t' j(| d| _)W | j" D ]\}\}}t%t|| qn(| j" D ]\}\}}t%t|| q:0 | j)j*D ]}|j+dkr|j,|v rd|_-tj.|_n\|g}t/0 }|r|1d}d||< |t2|j3 7 }qt4| D ]}| j)5| q|j+dkr`d|_q`| j)S )a  
        Traces `root` and returns the corresponding FX `torch.fx.Graph` representation. `root` can either be a
        `torch.nn.Module` instance or a Python callable. Note that after this call, `self.root` may be different from
        the `root` passed in here. For example, when a free function is passed to `trace()`, we will create a
        `torch.nn.Module` instance to use as the root and add embedded constants to.

        Args:
            root (`torch.nn.Module` or  `Callable`):
                Either a `torch.nn.Module`` or a function to be traced through. If root is not a
                [`~transformers.PreTrainedModel`], then `dummy_inputs` must be passed, otherwise tracing will fail.
            concrete_args (`Dict[str, Any], *optional*):
                Concrete arguments that should not be treated as Proxies
            dummy_inputs (`Dict[str, Any]`, *optional*):
                The dummy inputs needed to handle data-dependent control-flow if `root` is not a
                [`~transformers.PreTrainedModel`]. It can also be used when `root` is a
                [`~transformers.PreTrainedModel`] to specify custom dummy inputs for a subset or all the model inputs.
            complete_concrete_args_with_inputs_not_in_dummy_inputs (`bool`, *optional*, defaults to `True`):
                If `True`, and `dummy_inputs` is specified, every argument that `root` can take that is not in
                `dummy_inputs` and not in `concrete_args` will be added to `concrete_args`, otherwise does nothing.

        Returns:
            `torch.fx.Graph`:
                A FX `torch.fx.Graph` representing the semantics of the passed-in `root`.

        Nz6You need to specify a default value for the parameter rI  c                    s*   i | ]"}|j vr|j  vr|j |jqS r:   rU  r3   rx   r   re  rf  r:   r;   
<dictcomp>  s   z"HFTracer.trace.<locals>.<dictcomp>r      r9  rf   Z_deserialize_graph_modulezCould not generate input named z8 for because root is not a transformers.PreTrainedModel.c                 S   s,   i | ]$\}}|t |tjr$|d n|qS )rC   )r5   rF   rs   r_   )rx   r,  Zinput_r:   r:   r;   rk  .  s   z**c                 S   s   i | ]}|t tt|qS r:   )r  r  rF   )rx   r   r:   r:   r;   rk  6  s   re  rE  r:   r   r   )6r]  r^  r5   rF   r	   Modulerd  r_  r   rU  r3   rb  rG   r[   updater4   r!  r   r   r   r   r   rt   supported_archsrN  r   
startswithrD  r  itemsrT  VAR_KEYWORDrJ  _TORCH_METHODS_TO_PATCHZpatched_torch_methodssetrK  setattraddr   tracegraphnodesopr   rl   rs   r?  OrderedDictru   rr   ZusersreversedZ
erase_node)rK   r  re  rf  rg  sigparaminput_namesr   Zsequence_lengthrH   Znum_choicesr8  r,  Zconcrete_metasrU  r  origr   r  Zto_visitZ	to_deletern   userr   rj  r;   rx    s     





zHFTracer.trace)rY  r2   c                 C   s   t dd |j D S )z
        Whether the module was instantiated with Proxies. If that is the case, such module cannot be a leaf module
        because its attributes are input-dependent.
        c                 s   s   | ]}t |tV  qd S rQ   r  )rx   r  r:   r:   r;   r~   g  r{   zKHFTracer._stateless_mod_instanciation_depends_on_proxies.<locals>.<genexpr>)any__dict__r   )rK   rY  r:   r:   r;   /_stateless_mod_instanciation_depends_on_proxiesb  s    z8HFTracer._stateless_mod_instanciation_depends_on_proxiesc                 C   s   |  |rdS d}|jj }| d| }d}t| j|rjt| j||u rRd}qj| d| }|d7 }q0|s|| j|| |S )zb
        Helper method which tries to insert a module that was not declared as submodule.
        r/   r   r   FTrf   )r  r   r   lowerr   r  r  Z
add_module)rK   rY  idxmod_namepathZalready_insertedr:   r:   r;   _insert_module_as_submodulei  s    

z$HFTracer._insert_module_as_submodulec              
      s   zt  |W S  ty~ } zX| jrftt| dkrftt| dkrf| |}|W  Y d}~S |W Y d}~n
d}~0 0 dS )ag  
        Helper method to find the qualified name of `mod` in the Module hierarchy of `root`. For example, if `root` has
        a submodule named `foo`, which has a submodule named `bar`, passing `bar` into this function will return the
        string "foo.bar".

        Args:
            mod (str): The `Module` to retrieve the qualified name for.
        r   N)	r   path_of_module	NameErrorr$  ri   rr   r_  buffersr  )rK   rY  r[  r  r   r:   r;   r    s    	.
zHFTracer.path_of_module)r   module_qualified_namer2   c                    s   |  | ot ||S rQ   )r  r   is_leaf_module)rK   r   r  r   r:   r;   r    s    zHFTracer.is_leaf_module)Zis_backward_compatibler   )objr2   c                 C   s"   t |d }|jjdkr|jS |S )zCalled when a proxy object is has the keys() method called.
        This is what happens when ** is called on a proxy. This should return an iterator if ** is supposed to work in
        your custom tracer.
        r4   z**kwargs)r   r  r   r   )rK   r  	attributer:   r:   r;   r4     s    zHFTracer.keys)NNN)NNT)*r   r   r   r   r#  bool__annotations__r$  rt  r,   r   r.   rp  r   r  r6   r   rk   r   rF   rs   rD  r   rc  r   r  rF  r  r   r	   rn  r   r   r
   rx  r  r  r  r  r   r4   r  r:   r:   r   r;   r"    s>   


D&    r"  )r+  r  c                    s|   t | j}t t|j ksdt dkr6 d nd }d|j }td| d|  fdd|j	 D S )Nrf   r   , z(The model does not have input(s) named: z&, expected a subset of the following: c                    s    i | ]}|j  vr|j |jqS r:   rh  ri  r  r:   r;   rk    r{   z%get_concrete_args.<locals>.<dictcomp>)
r]  r^  rd  ru  r_  r4   ri   joinr[   r   )r+  r  r~  Zformatted_input_namesZformatted_allowed_input_namesr:   r  r;   get_concrete_args  s    r  )r+  c                 C   s2   | j jtvr.dt}td| j j d| d S )Nr  zModel z) is not supported yet, supported models: )r   r   _SUPPORTED_MODELSr  r>  )r+  Zsupported_model_namesr:   r:   r;   check_if_model_is_supported  s
    
r  )r+  r  disable_check
tracer_clsr2   c                 C   sn   |du r| j  }t|}t| |}|s0t|  | }|j| |d}tj| |}| j	|_	| j
|_| j|_|S )a  
    Performs symbolic tracing on the model.

    Args:
        model ([`PretrainedModel`]):
            The model to trace.
        input_names (`List[str]`, *optional*):
            The names of the inputs of the traced model. If unset, model.dummy_inputs.keys() are used instead.
        disable_check (`bool`, *optional*, defaults to `False`):
            If `True`, no check is done before trying to trace the model, this is mostly usesul for debugging purposes.
        tracer_cls (`Type[HFTracer]`, *optional*, defaults to `HFTracer`):
            The tracer class to use for instantiating the tracer. If unset, `HFTracer` is used instead.

    Returns:
        `torch.fx.GraphModule`: A GraphModule constructed by recording operations seen while tracing the model.

    Example:

        ```python
        from transformers.utils.fx import symbolic_trace

        traced_model = symbolic_trace(model, input_names=["input_ids", "attention_mask", "token_type_ids"])
        ```
    Nrm  )rf  r4   rr   r  r  rx  rF   r
  r   r=  r   r-  rD   )r+  r  r  r  re  r   Ztraced_graphZtracedr:   r:   r;   symbolic_trace  s    

r  )N)NNrN   FF)F)NN)NN)N)N)N)rB   )r  r  N)r  r?  r  r]  r   r   osr  rR  typingr   r   r   r   r   r   r   rF   r	   Ztorch.fxr
   r   r   r   Ztorch.fx._compatibilityr   Ztorch.fx.proxyr   r/   r   r   r   Zmodels.autor   Zmodels.auto.modeling_autor   r   r   r   r   r   r   r   r   r   r    r!   r"   r#   r$   r%   r&   r'   r(   utilsr)   r*   r+   r,   r-   Zpeftr.   Z
get_loggerr   loggerenvironr7   upperrQ  r6   r<   Z(_REGULAR_SUPPORTED_MODEL_NAMES_AND_TASKSZ_REGULAR_SUPPORTED_MODELSitemr5   rt   extendZ_SPECIAL_SUPPORTED_MODELSr   sortedru  r  rM   rP   rR   rS   rU   rX   rY   r\   rb   re   rq   rw   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   Z	EmbeddingZ
functionalZ	embeddingZ	LayerNormZ	GroupNormZLinearZreluZReLUwhereabsr%  rv   catstackrw  mulrs   matmulZbmmZbaddbmmr   repeatZrepeat_interleaveZrollZflipZindex_selectZgatherZConv1dZConv2dZsqueezeZ	unsqueezer   Zone_hotZMSELossZCrossEntropyLossZBCEWithLogitsLossr   r   r  r   r   r   r  r  rk   r!  r"  rn  r  r  r  r  r:   r:   r:   r;   <module>   s$  
$T	
 (- 

	$






/,	   Y

