o
    ;ήcX                     @   s   d dl mZ d dlmZ d dlmZ d dlmZ d dlm	Z	m
Z
 d dl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mZmZ d dlmZ d dlmZmZ eeZG dd dZ G dd deZ!dS )    )Enum)	getLogger)name)path)TupleUnionN)Fusion)AttentionMaskFormat)FusionUtilsNumpyHelper)	NodeProtoTensorProtohelpernumpy_helper)	OnnxModel)SymbolicShapeInferenceHelperget_shape_from_type_protoc                   @   sN   e Zd ZdZdefddZdefddZdd	 Zd
d Z	de
de
fddZdS )AttentionMask:
    Fuse Attention subgraph into one Attention node.
    modelc                 C   s(   || _ i | _i | _t|| _tj| _d S N)r   mask_indicemask_castedr
   utilsr	   MaskIndexEndmask_format)selfr    r   P/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/fusion_attention.py__init__   s
   
zAttentionMask.__init__r   c                 C   s
   || _ d S r   )r   )r   r   r   r   r   set_mask_format$   s   
zAttentionMask.set_mask_formatc                 C   s*   || j v r|| j | ksJ || j |< d S r   )r   )r   mask
mask_indexr   r   r   set_mask_indice'   s   
zAttentionMask.set_mask_indicec                 C   s    t | jdks	J tt| jS )Nr   )lenr   nextiter)r   r   r   r   get_first_mask,   s   zAttentionMask.get_first_maskinputreturnc                 C   s   | j tjkrd S || jv r| j| S | j|r!| j|\}}n
| j|\}}d}|r2|| j	|< | j tj
kr?|| j|< |S | jd}tjd|g|g| jddd}|jtddgtdd	g | j| || j|< |S )
NTr"   	ReduceSumMaskReduceSuminputsoutputsr   axes   keepdimsr   )r   r	   NoMaskr   r   find_graph_inputr   cast_graph_input_to_int32cast_input_to_int32r   r   create_node_namer   	make_node	attributeextendmake_attributeadd_node)r   r(   casted
input_name	cast_nodeoutput_namemask_index_noder   r   r   process_mask0   s0   



"
zAttentionMask.process_maskN)__name__
__module____qualname____doc__r   r   r	   r    r#   r'   strrA   r   r   r   r   r      s    	r   c                       s   e Zd ZdZdedededef fddZded	e	eef fd
dZ
defddZdedededededededededededed	eedf fddZdd Z  ZS )FusionAttentionr   r   hidden_size	num_headsattention_maskc                    s6   t  |dddg || _|| _|| _d| _d| _d S )N	AttentionSkipLayerNormalizationLayerNormalizationT)superr   rH   rI   rJ   num_heads_warninghidden_size_warning)r   r   rH   rI   rJ   	__class__r   r   r   Z   s   
zFusionAttention.__init__	reshape_qr)   c                 C   s  | j |jd }|du rt|jd  d | j| jfS t|}t	|dks5|d dks5|d dkrDtd| d	 | j| jfS |d }|d }|| }| jdkrm|| jkrm| j
rmtd
| j d| d d| _
| jdkr|| jkr| jrtd| j d| d d| _||fS )zDetect num_heads and hidden_size from a reshape node.

        Args:
            reshape_q (NodeProto): reshape node for Q

        Returns:
            Tuple[int, int]: num_heads and hidden_size
        r0   Nz is not initializer.      r      zq_shape_value=z7. Expected value are like [0, 0, num_heads, head_size].z--num_heads is z. Detected value is z. Using detected value.Fz--hidden_size is )r   get_initializerr(   loggerdebugrI   rH   r   to_arrayr$   rO   warningrP   )r   rS   q_shapeq_shape_valuerI   	head_sizerH   r   r   r   get_num_heads_and_hidden_sizej   s,   
$z-FusionAttention.get_num_heads_and_hidden_sizeadd_qkc                 C   s   | j jdd}|d u rd S ||jd }||jd }|d u s%|d u r0td| d d S ||kr?td| d d S |jd S )	NT)updater   r0   zone of the inputs of z is Nonezthe shape of two inputs of z is not same)r   infer_runtime_shapeget_edge_shaper(   rX   rY   )r   r`   shape_inferinput_0_shapeinput_1_shaper   r   r   get_add_qk_str   s   
zFusionAttention.get_add_qk_strr"   q_matmulk_matmulv_matmulq_addk_addv_addr(   output
add_qk_strNc           ,      C   sd  |dksJ |	dkr|	| dkrt d|	 d|  dS | j|jd }| j|jd }| j|jd }| j|jd pI| j|jd }| j|jd p[| j|jd }| j|jd pm| j|jd }|du r~t|jd  d dS |r|r|r|sdS t|}t|}t|}|j|jksJ |jd }|jd }|jd }||  kr|ksJ  J |	dkr|	|krt 	d|	 d| d	 d
}|j|jkrd}t
|jdd }t
|jdd }t
|jdd }d}|rt
j|||fdd}|| | }nt
j|||fdd}d| }t|}t|} t|}!t
|j}"t
| j}#t
|!j}$|"|#  krJ|ksMJ  J |$|ksTJ d}%|rjt
j|| |!fdd}&|"|# |$ }%nt
j|| |!fdd}&d|" }%| jd}'tj|'d tj||g|  d}(|jdkr|(tt|(t
j|(j | j|(| j tj|'d tj|%g|&  d})|jdkr|)tt|)t
j|)j | j|)| j |
|'d |'d g}*|dur|*| n|*d |dur|*d |*| tjd|*|g|'d}+d|+_ |+j!"t#d|g |r0|+j!"t#d|||gg |+S )a  Create an Attention node.

        Args:
            mask_index (str): mask input
            q_matmul (NodeProto): MatMul node in fully connection for Q
            k_matmul (NodeProto): MatMul node in fully connection for  K
            v_matmul (NodeProto): MatMul node in fully connection for  V
            q_add (NodeProto): Add bias node in fully connection for Q
            k_add (NodeProto): Add bias node in fully connection for K
            v_add (NodeProto): Add bias node in fully connection for V
            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
            hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
            input (str): input name
            output (str): output name

        Returns:
            Union[NodeProto, None]: the node created or None if failed.
        r   zinput hidden size z# is not a multiple of num of heads Nr0   zl is not an initializer. Please set do_constant_folding=True in torch.onnx.export to unblock attention fusionzInput hidden size (z3) is not same as weight matrix dimension of q,k,v (z:). Please provide a correct input hidden size or pass in 0FT)axisrV   rK   _qkv_weight)r   	data_typedimsvals
   	_qkv_bias r,   zcom.microsoftrI   qkv_hidden_sizes)$rX   rY   r   rW   r(   printr   rZ   shaper[   npprodconcatenatestackr6   r   make_tensorr   FLOATflattentolistrr   CopyFromr   
from_arrayastypefloat16r   add_initializerthis_graph_nameappendr7   domainr8   r9   r:   ),r   r"   rh   ri   rj   rk   rl   rm   rI   rH   r(   rn   ro   q_weightk_weightv_weightq_biask_biasv_biasqwkwvw
qw_in_size
kw_in_size
vw_in_sizeis_qkv_diff_dimsqw_out_sizekw_out_sizevw_out_sizeqkv_weight_dim
qkv_weightqbkbvbq_bias_shapek_bias_shapev_bias_shapeqkv_bias_dimqkv_biasattention_node_nameweightbiasattention_inputsattention_noder   r   r   create_attention_node   s   !$$$








 
"
"




z%FusionAttention.create_attention_nodec           3      C   sp  |}|j dkr| j|dd}|d ur|}nd S | j|g dg d}d }|d ur2|\}}}	}
}n| j|g dg d}|d urI|\}}}
}nd S g }t|jD ]\}}||vr[qR||d jd kreqR|| qRt|dkrsd S |d }	 | j|d	d}|d ur||jd  }|d urt|d
kr|d }|j dkr|jd }n,d S |d urt|dkr|jd }nd S |j dkr|| }|D ]}|j dkr|jd }q|| }dd |D }|	ddkrd S | j|g dg d}|d u rt
d d S |\}}}}d}d}g dg dfg dg dfg dg dfg dg dfd}d }| D ]%\}}| j||d |d }|d u rAq*|dkrHd}|dkrOd} |d u r\t
d d S d }d } d }!|rl|\}}!} }n|rv|\}}}!} n|\}}}} | j| g dg d }"|"d u r| j| g d!g d"}"|"d u rt
d# d S |"d$ }#|"d% }$|"d& }%| j| g dg d}&|&d u r| j| g d'g d(}&|&d u rt
d) d S |&d% }'|&d& }(d })d }*|r
| j|!g d*g d+fg d,g d+fg d-g d.fg|\}})}nO|r@| j|!g d/g d.fg d,g d+fg|\}})}|d ur?| |}*|*d u r?t
d0|  d S n| j|g d1g d2fg d3g d4fg|\}})}|)d u ret
d5 d S |jd |kr2|%jd |kr4|(jd |kr6| j|)d& jd }+|d u r|	n|
},| |#\}-}.| |+|%|(||$|'||-|.||,jd |*}/|/d u rd S | j|/ | j| j|/j< |d ur|jd }0d6|0 }1tjd7|0 tjd8gtdd|-t|.|- g dd9}2| j|2| j | j t!d:|,jd |2jg|1gd;|0 | j |1|jd< | j"#|,|
|g | j"#| | j"#|" | j"#|& | j"#| d| _$d S d S d S d S )<NrM   Addr   )r   MatMulReshape	Transposer   )NNr   r   r   )r   Einsumr   r   )r0   Nr   r   r0   MulrU      c                 S   s   g | ]}|j qS r   )op_type).0childr   r   r   
<listcomp>  s    z(FusionAttention.fuse.<locals>.<listcomp>r   rV   )r   r   r   r   )r0   r   r   Nz&fuse_attention: failed to match v pathF)Softmaxr   Divr   )r   r   Nr   )r   r   r   r   )r   Wherer   r   )r   r   rU   r   )r   r   r   r   )r   r   r   rU   )path1path2path3path4r   Tr   z'fuse_attention: failed to match qk path)r   r   r   N)r   r   r   r   r   )r   r   r   r   Nz&fuse_attention: failed to match q path)r   r   r   r   r   )r0   r   r   r   Nz&fuse_attention: failed to match k path)Expandr   Equal)r   r   r   )r   	Unsqueezer   )Castr   r   r   )r   r   r   r   )r   r   r   r   z4fuse_attention: failed to verify shape inference of )r   Subr   r   r   )Nr   r0   r   r   )r   r   r   r   )Nr   r0   r   z)fuse_attention: failed to match mask pathedge_modified_shape_modified_tensorrT   )r   rr   rs   rt   rawr   reshape_modified_)%r   r   match_parentmatch_parent_path	enumerater(   rn   r   r$   countrX   rY   itemsmatch_parent_pathsrg   rJ   rA   r_   r   nodes_to_addr   node_name_to_graph_namer   r   r   r   INT64r{   int64inttobytesr   r;   r7   nodes_to_remover9   prune_graph)3r   normalize_nodeinput_name_to_nodesoutput_name_to_node
start_nodeadd_before_layernorm	qkv_nodeseinsum_node_reshape_qkvtranspose_qkv
matmul_qkvother_inputsir(   
root_inputmul_before_layernormmul_childrenlayernorm_nodechildrenr   children_typesv_nodesadd_vmatmul_v
is_distillis_distill_addqk_pathsqk_nodeskvr`   	matmul_qkwhere_qkq_nodesrS   add_qmatmul_qk_nodesadd_kmatmul_k
mask_nodesro   r"   attention_last_nodeq_num_headsq_hidden_sizenew_nodeunique_indexnew_edgeshape_tensorr   r   r   fuseI  s  
















	






0



	
zFusionAttention.fuse)rB   rC   rD   rE   r   r   r   r   r   r   r_   rg   rF   r   r   r  __classcell__r   r   rQ   r   rG   U   sT    '	


 'rG   )"enumr   loggingr   osr   sysr   typingr   r   numpyr{   fusion_baser   fusion_optionsr	   fusion_utilsr
   r   onnxr   r   r   r   
onnx_modelr   shape_infer_helperr   r   rB   rX   r   rG   r   r   r   r   <module>   s   ?