o
    ;ήcF                     @   sV   d dl 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 ZG dd dZdS )	    )PathN   )attribute_to_kwargfind_by_namec              	   C   s  t  }|dd | jD  |dd | jD  g }| jD ]w}|}dd |jD }|ri }|jD ]M}i }	|jtjjkrOt	|j
|\}
}|j|
i}	|| n*|jtjjkrug }|jD ]}t	||\}
}||
 || q[|j|i}	nt|}	||	 q1tj|j|j|jfd|ji|}|| q| d | j| |dd | jD  g }| jD ]}|j|v r||j q|| qd	d
 | jD }|D ]/}| j| |j|v rz| j||j  W q ty   |jdk rtd|j Y qw q|dd | jD  | |fS )zClean unused initializers from graph.

    Returns:
        A cleaned graph without unused initializers
        A list of tensor names, which are not produced by this graph and its subgraphes
    c                 s   s$    | ]}|j D ]}|r|V  qqd S N)input).0node
input_name r   J/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/quantization/onnx_model.py	<genexpr>   s   " z-_clean_initializers_helper.<locals>.<genexpr>c                 s   s    | ]	}|j r|j V  qd S r   name)r   g_outr   r   r   r      s    c                 S   s,   g | ]}|j tjjks|j tjjkr|qS r   )typeonnxAttributeProtoGRAPHGRAPHSr   attrr   r   r   
<listcomp>   s    z._clean_initializers_helper.<locals>.<listcomp>r   r	   c                 s   s     | ]}|j D ]}|V  qqd S r   )output)r   r	   r   r   r   r   r   ;   s    c                 S   s   i | ]}|j |qS r   r   r   r   r   r   r   
<dictcomp>E   s    z._clean_initializers_helper.<locals>.<dictcomp>   zFWarning: invalid weight name {} found in the graph (not a graph input)c                 s       | ]}|j V  qd S r   r   r   r   r   r   r   S       )setupdater	   r   	attributer   r   r   r   _clean_initializers_helpergr   r   graphsappendr   onnx_helper	make_nodeop_typer   
ClearFieldextenddifference_updateinitializerremoveStopIteration
ir_versionprintformat)graphmodelrequesting_tensor_names	new_nodesr	   new_nodegraph_attrskwargsr   new_attributecleaned_sub_graphsub_requesting_tensor_namescleaned_graphessubgraphunused_initializerr,   name_to_inputr   r   r   r"   
   sx   





"




r"   c                   @   sN  e Zd Zdd Zdd Zdd Zdd Zd	d
 Zdd Zdd Z	dd Z
dd Zdd Zdd Zdd Zdd Zdd Zdd Zdd  Zd!d" Zd#d$ ZdMd&d'ZdMd(d)ZdMd*d+Zd,d- Zd.d/ Zed0d1 Zed2d3 Zd4d5 ZdNd7d8Zed9d: Zd;d< Z ed=d> Z!d?d@ Z"dAdB Z#dCdD Z$dEe%dFe&fdGdHZ'dIdJ Z(dKdL Z)d%S )O	ONNXModelc                 C   s
   || _ d S r   )r3   )selfr3   r   r   r   __init__Y      
zONNXModel.__init__c                 C   
   | j jjS r   )r3   r2   r	   rA   r   r   r   nodes\   rC   zONNXModel.nodesc                 C   rD   r   )r3   r2   r,   rE   r   r   r   r,   _   rC   zONNXModel.initializerc                 C      | j jS r   )r3   r2   rE   r   r   r   r2   b      zONNXModel.graphc                 C   rG   r   )r3   r/   rE   r   r   r   r/   e   rH   zONNXModel.ir_versionc                 C   rG   r   )r3   opset_importrE   r   r   r   rI   h   rH   zONNXModel.opset_importc                 C   s&   || j jjv r| j jj| d S d S r   )r3   r2   r	   r-   rA   r	   r   r   r   remove_nodek   s   zONNXModel.remove_nodec                 C      |D ]}|  | qd S r   )rK   )rA   nodes_to_remover	   r   r   r   remove_nodeso      zONNXModel.remove_nodesc                 C   s   | j jj|g d S r   r3   r2   r	   r*   rJ   r   r   r   add_nodes   s   zONNXModel.add_nodec                 C   s   | j jj| d S r   rP   )rA   nodes_to_addr   r   r   	add_nodesv   s   zONNXModel.add_nodesc                 C   s0   t |j| jjjd u r| jjj|g d S d S r   )r   r   r3   r2   r,   r*   )rA   tensorr   r   r   add_initializery   s   zONNXModel.add_initializerc                 C   s&   | j jjD ]}|j|kr|  S qd S r   )r3   r2   r,   r   )rA   r   rT   r   r   r   get_initializer}   s
   
zONNXModel.get_initializerc                 C   s   t dd | jjjD S )Nc                 s   r   r   r   )r   r,   r   r   r   r      r   z5ONNXModel.get_initializer_name_set.<locals>.<genexpr>)r   r3   r2   r,   rE   r   r   r   get_initializer_name_set   s   z"ONNXModel.get_initializer_name_setc                 C   sX   || j jjv r(| j jj| | j jjD ]}|j|jkr'| j jj|  d S qd S d S r   )r3   r2   r,   r-   r   r   )rA   rT   r   r   r   r   remove_initializer   s   zONNXModel.remove_initializerc                 C   rL   r   )rX   )rA   init_to_remover,   r   r   r   remove_initializers   rO   zONNXModel.remove_initializersc                 C   s8   |   }t }| jjjD ]}|j|vr||j q|S r   )rW   r   r3   r2   r   r   add)rA   initializer_namesnon_initializer_inputsr   r   r   r   get_non_initializer_inputs   s   
z$ONNXModel.get_non_initializer_inputsc                 C   sF   i }| j jjD ]}|jD ]}||vr|g||< q|| | qq|S r   )r3   r2   r	   r   r%   )rA   input_name_to_nodesr	   r
   r   r   r   r_      s   
zONNXModel.input_name_to_nodesc                 C   s,   i }| j jjD ]}|jD ]}|||< qq|S r   )r3   r2   r	   r   )rA   output_name_to_noder	   output_namer   r   r   r`      s   

zONNXModel.output_name_to_nodeNc                 C   sD   |d u r|   }g }|jD ]}||v r|| D ]}|| qq|S r   )r_   r   r%   )rA   r	   r_   childrenr   r   r   r   get_children   s   
zONNXModel.get_childrenc                 C   s:   |d u r|   }g }|jD ]}||v r|||  q|S r   )r`   r   r%   )rA   r	   r`   parentsr   r   r   r   get_parents   s   
zONNXModel.get_parentsc                 C   s@   |d u r|   }t|j|krd S |j| }||vrd S || S r   )r`   lenr   )rA   r	   idxr`   r   r   r   r   
get_parent   s   
zONNXModel.get_parentc                 C   s"   t |j}|| t||}|S )zFind out if a node exists in a graph or a node is in the
        new set of nodes created during quantization.

        Returns:
            The node found or None.
        )listr	   r*   r   )rA   	node_namenew_nodes_listr2   graph_nodes_listr	   r   r   r   find_node_by_name   s   


zONNXModel.find_node_by_namec                 C   s4   g }|j D ]}|jD ]}||jkr|| q
q|S )zD
        Find all nodes with given initializer as an input.
        )r	   r   r   r%   )rA   r2   r,   rF   r	   
node_inputr   r   r   find_nodes_by_initializer   s   



z#ONNXModel.find_nodes_by_initializerc                 C   sL   t t|d ddD ]}|| }|jD ]}|j| kr"||f    S qq
dS )Nr   )NN)rangerf   r,   r   )r   
graph_pathgidr2   rT   r   r   r   __get_initializer   s   

zONNXModel.__get_initializerc                 C   s>  g }| d }|j D ]}dd |jD }t|roi }|jD ]@}|jdkr3| |j |jt| i}n%|jdkrTg }|j	D ]}	| |	 |
t| g q=|j|i}nt|}|| qtj|j|j|jfd|ji|}|jdkrd}
d}d	}d	}|jD ]-}|jd
krt|}
q|jdkrt|}q|jdkrt|}q|jdkrt|}q|
dkr|dkr|d	kr|jd }|dkr't|jd | \}}|rt|}t|j}|j|_|j| |jD ]}|j|kr|j|  nq|j
|g n"|d7 }tjd|jd g|g|jdkr|jd ndd}|| tjd|jd	 |g|jd	 t|jdkr>dnd g|jdkrL|jd ndd}|| t|jdkrtjd|jd	 d |jd g|j|jdkrx|jd ndd}|| q	|| q	|| q	|d |j 
| |   |S )Nrp   c                 S   s$   g | ]}|j d ks|j dkr|qS )   
   )r   r   r   r   r   r      s   $ z8ONNXModel.__replace_gemm_with_matmul.<locals>.<listcomp>ru   rv   r   Gemmg      ?r   alphabetatransAtransBr   _Transposed	Transpose 
_Transpose)inputsoutputsr   MatMul   _MatMulAdd_Addr	   )r	   r!   rf   r   r%   r#   r   r@   $_ONNXModel__replace_gemm_with_matmulr$   r*   r   r    r&   r'   r(   r   r   get_attribute_value_ONNXModel__get_initializeronnx_numpy_helperto_array
from_arrayTr,   r-   r)   pop)rr   r5   r2   r	   r7   r8   r   kvvaluer=   rx   ry   rz   r{   inputBBBs_graphB_arrayB_transr   transpose_nodematmul_noderQ   r   r   r   __replace_gemm_with_matmul   s   




"












"


z$ONNXModel.__replace_gemm_with_matmulc                 C   s   |   g}t| d S r   )r2   r@   r   )rA   rr   r   r   r   replace_gemm_with_matmulI  s   
z"ONNXModel.replace_gemm_with_matmulFc                 C   s<   |    |rtjj| jdt|jd d t| j| dS )zS
        Save model to external data, which is needed for model size > 2GB
        Tz.data)all_tensors_to_one_filelocationN)topological_sortr   external_data_helperconvert_model_to_external_datar3   r   r   
save_model)rA   output_pathuse_external_data_formatr   r   r   save_model_to_fileM  s   zONNXModel.save_model_to_filec                 C   H   t |tr
t |tsJ tt| jD ]}| j| |kr!|| j|< qd S r   )
isinstancestrrq   rf   r   )r	   old_input_namenew_input_namejr   r   r   replace_node_inputZ     
zONNXModel.replace_node_inputc                 C   "   | j jjD ]	}t||| qd S r   )r3   r2   r	   r@   r   )rA   r   r   r	   r   r   r   replace_input_of_all_nodesa     z$ONNXModel.replace_input_of_all_nodesc                 C   r   r   )r   r   rq   rf   r   )r	   old_output_namenew_output_namer   r   r   r   replace_node_outpute  r   zONNXModel.replace_node_outputc                 C   r   r   )r3   r2   r	   r@   r   )rA   r   r   r	   r   r   r   replace_output_of_all_nodesl  r   z%ONNXModel.replace_output_of_all_nodesc                 C   s   |   }g }|  }|D ]}|jdkr'| |jd s'|jd |vr'|| q| | g }|  D ](}|j|vr[| |js[|| | 	 j
D ]}|j|jkrZ| 	 j
| qJq3| | d S )NConstantr   )r_   rF   r(   is_graph_outputr   r%   rN   r,   r   r2   r   r-   rZ   )rA   r_   unused_nodesrF   r	   ununsed_weightswgraph_inputr   r   r   remove_unused_constantp  s,   


z ONNXModel.remove_unused_constantc                 C   $   | j jjD ]
}|j|kr dS qdS NTF)r3   r2   r   r   )rA   ra   r   r   r   r   r     
   
zONNXModel.is_graph_outputtensor_namereturnc                 C   r   r   )r3   r2   r   r   )rA   r   r   r   r   r   is_graph_input  r   zONNXModel.is_graph_inputc                 C   s  dgt |   }i }g }t|  D ]7\}}tdd |jD ||< || dkr3||  |  q|jD ]}||vrB|g||< q6|| | q6qdd |  D }dd | jjjD }|| }	|		  d }
|	D ]+}|
|krqqj|}
||v r|| D ]}|| d ||< || dkr||  |  q{qjd}t |}||k r|| j
D ](}||v r|| D ]}|| d ||< || dkr||  |  |d }qq|d }||k s|t |  jksJ d|  d	 |  j| d S )
Nr   c                 s   s    | ]}|rd V  qdS )r   Nr   )r   _r   r   r   r     s    z-ONNXModel.topological_sort.<locals>.<genexpr>c                 S      g | ]}|j qS r   r   )r   initr   r   r   r         z.ONNXModel.topological_sort.<locals>.<listcomp>c                 S   r   r   r   r   r   r   r   r     r   r   zGraph is not a DAGr	   )rf   rF   	enumeratesumr   r%   r,   r3   r2   sortr   r	   r)   r*   )rA   
deps_countdeps_to_nodessorted_nodesnode_idxr	   r
   r\   graph_input_namesinput_namesprev_input_namestartendr   r   r   r   r     sX   

zONNXModel.topological_sortc                 C   s   t |  | jS r   )r"   r2   r3   rE   r   r   r   clean_initializers  s   zONNXModel.clean_initializersr   )F)*__name__
__module____qualname__rB   rF   r,   r2   r/   rI   rK   rN   rQ   rS   rU   rV   rW   rX   rZ   r^   r_   r`   rc   re   rh   rm   ro   staticmethodr   r   r   r   r   r   r   r   r   r   r   boolr   r   r   r   r   r   r   r@   X   sR    






[


2r@   )pathlibr   r   onnx.helperhelperr&   onnx.numpy_helpernumpy_helperr   quant_utilsr   r   r"   r@   r   r   r   r   <module>   s    N