Ë
    ˆ\;j†-  ã                   ó¨  — d dl Z 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 g Z	dd„Z
d„ Zd„ Zd	„ Zd
„ Zd„ Zd„ Zd„ Zd„ Zd„ Zi ej(                  e“ej*                  e“ej,                  e“ej.                  e“ej0                  e“ej2                  e“ej4                  j6                  j8                  e“ej:                  e“ej<                  e“ej>                  e“ej@                  e“ejB                  e“ejD                  e“ejF                  e“ejH                  e“ejJ                  e“ejL                  e“ejN                  eejP                  ei¥Z)dd„Z*y)é    N)Únn)Úunwrap_decoratorsé   )ÚTableÚstatic_flopsc                 óF  — t        | t        j                  «      rAt        | j                  «      \  }| _        t        j                  |«      }t        | |||¬«      S t        | t
        j                  j                  «      rt        | |¬«      S t        j                  d«       y)a  Print a table about the FLOPs of network.

    Args:
        net (paddle.nn.Layer||paddle.static.Program): The network which could be a instance of paddle.nn.Layer in
                    dygraph or paddle.static.Program in static graph.
        input_size (list): size of input tensor. Note that the batch_size in argument ``input_size`` only support 1.
        custom_ops (A dict of function, optional): A dictionary which key is the class of specific operation such as
                    paddle.nn.Conv2D and the value is the function used to count the FLOPs of this operation. This
                    argument only work when argument ``net`` is an instance of paddle.nn.Layer. The details could be found
                    in following example code. Default is None.
        print_detail (bool, optional): Whether to print the detail information, like FLOPs per layer, about the net FLOPs.
                    Default is False.

    Returns:
        Int: A number about the FLOPs of total network.

    Examples:
        .. code-block:: python

            >>> import paddle
            >>> import paddle.nn as nn

            >>> class LeNet(nn.Layer):
            ...     def __init__(self, num_classes=10):
            ...         super().__init__()
            ...         self.num_classes = num_classes
            ...         self.features = nn.Sequential(
            ...             nn.Conv2D(1, 6, 3, stride=1, padding=1),
            ...             nn.ReLU(),
            ...             nn.MaxPool2D(2, 2),
            ...             nn.Conv2D(6, 16, 5, stride=1, padding=0),
            ...             nn.ReLU(),
            ...             nn.MaxPool2D(2, 2))
            ...
            ...         if num_classes > 0:
            ...             self.fc = nn.Sequential(
            ...                 nn.Linear(400, 120),
            ...                 nn.Linear(120, 84),
            ...                 nn.Linear(84, 10))
            ...
            ...     def forward(self, inputs):
            ...         x = self.features(inputs)
            ...
            ...         if self.num_classes > 0:
            ...             x = paddle.flatten(x, 1)
            ...             x = self.fc(x)
            ...         return x
            ...
            >>> lenet = LeNet()
            >>> # m is the instance of nn.Layer, x is the intput of layer, y is the output of layer.
            >>> def count_leaky_relu(m, x, y):
            ...     x = x[0]
            ...     nelements = x.numel()
            ...     m.total_ops += int(nelements)
            ...
            >>> FLOPs = paddle.flops(lenet,
            ...                      [1, 1, 28, 28],
            ...                      custom_ops= {nn.LeakyReLU: count_leaky_relu},
            ...                      print_detail=True)
            >>> print(FLOPs)
            <class 'paddle.nn.layer.conv.Conv2D'>'s flops has been counted
            <class 'paddle.nn.layer.activation.ReLU'>'s flops has been counted
            Cannot find suitable count function for <class 'paddle.nn.layer.pooling.MaxPool2D'>. Treat it as zero FLOPs.
            <class 'paddle.nn.layer.common.Linear'>'s flops has been counted
            +--------------+-----------------+-----------------+--------+--------+
            |  Layer Name  |   Input Shape   |   Output Shape  | Params | Flops  |
            +--------------+-----------------+-----------------+--------+--------+
            |   conv2d_0   |  [1, 1, 28, 28] |  [1, 6, 28, 28] |   60   | 47040  |
            |   re_lu_0    |  [1, 6, 28, 28] |  [1, 6, 28, 28] |   0    |   0    |
            | max_pool2d_0 |  [1, 6, 28, 28] |  [1, 6, 14, 14] |   0    |   0    |
            |   conv2d_1   |  [1, 6, 14, 14] | [1, 16, 10, 10] |  2416  | 241600 |
            |   re_lu_1    | [1, 16, 10, 10] | [1, 16, 10, 10] |   0    |   0    |
            | max_pool2d_1 | [1, 16, 10, 10] |  [1, 16, 5, 5]  |   0    |   0    |
            |   linear_0   |     [1, 400]    |     [1, 120]    | 48120  | 48000  |
            |   linear_1   |     [1, 120]    |     [1, 84]     | 10164  | 10080  |
            |   linear_2   |     [1, 84]     |     [1, 10]     |  850   |  840   |
            +--------------+-----------------+-----------------+--------+--------+
            Total Flops: 347560     Total Params: 61610
            347560
    )ÚinputsÚ
custom_opsÚprint_detail)r   zKYour model must be an instance of paddle.nn.Layer or paddle.static.Program.éÿÿÿÿ)Ú
isinstancer   ÚLayerr   ÚforwardÚpaddleÚrandnÚdynamic_flopsÚstaticÚProgramr   ÚwarningsÚwarn)ÚnetÚ
input_sizer
   r   Ú_r	   s         úbG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/hapi/dynamic_flops.pyÚflopsr      s„   € ôb �#”r—x‘xÔ ô +¨3¯;©;Ó7‰ˆˆ3Œ;ä—‘˜jÓ)ˆÜØ˜¨:ÀLô
ð 	
ô 
�CœŸ™×.Ñ.Ô	/Ü˜C¨lÔ;Ð;ä�‰ØYô	
ð ó    c                 óL  — |d   }t        j                  | j                  j                  dd  «      }| j                  �dnd}t        |j                  «       «      |j                  d   | j                  z  |z  |z   z  }| xj                  t        t        |«      «      z  c_        y ©Nr   é   r   )
ÚnpÚprodÚweightÚshapeÚbiasÚintÚnumelÚ_groupsÚ	total_opsÚabs)ÚmÚxÚyÚ
kernel_opsÚbias_opsr(   s         r   Úcount_convNdr/      sƒ   € Ø	ˆ!‰€AÜ—‘˜Ÿ™Ÿ™¨¨Ð+Ó,€JØ—F‘FÐ&‰q¨A€HÜ�A—G‘G“I“Ø	�‰�‰
�Q—Y‘YÑ Ñ+¨hÑ6ñ€Ið ‡K‚K”3”s˜9“~Ó&Ñ&†Kr   c                 ój   — |d   }|j                  «       }| xj                  t        |«      z  c_        y ©Nr   ©r&   r(   r%   )r*   r+   r,   Ú	nelementss       r   Úcount_leaky_relur4   ‰   s(   € Ø	ˆ!‰€AØ—‘“	€IØ‡K‚K”3�y“>Ñ!†Kr   c                 óž   — |d   }|j                  «       }| j                  sd|z  }| xj                  t        t	        «      «      z  c_        y )Nr   r   )r&   Útrainingr(   r)   r%   )r*   r+   r,   r3   r(   s        r   Úcount_bnr7   �   s=   € Ø	ˆ!‰€AØ—‘“	€IØ�:Š:Ø˜	‘Mˆ	Ø‡K‚K”3”s˜9“~Ó&Ñ&†Kr   c                 ó®   — | j                   j                  d   }|j                  «       }||z  }| xj                  t	        t        |«      «      z  c_        y r1   )r"   r#   r&   r(   r)   r%   )r*   r+   r,   Ú	total_mulÚnum_elementsr(   s         r   Úcount_linearr;   —   s@   € Ø—‘—‘˜qÑ!€IØ—7‘7“9€LØ˜LÑ(€IØ‡K‚K”3”s˜9“~Ó&Ñ&†Kr   c                 ón   — d}|j                  «       }||z  }| xj                  t        |«      z  c_        y )Nr   r2   )r*   r+   r,   r-   r:   r(   s         r   Úcount_avgpoolr=   ž   s.   € Ø€JØ—7‘7“9€LØ˜\Ñ)€Ià‡K‚K”3�y“>Ñ!†Kr   c                 óD  — t        j                  |d   j                  dd  «      t        j                  |j                  dd  «      z  }t        j                  |«      }d}||z   }|j	                  «       }||z  }| xj
                  t        t        |«      «      z  c_        y r   )r    Úarrayr#   r!   r&   r(   r)   r%   )	r*   r+   r,   ÚkernelÚ	total_addÚ	total_divr-   r:   r(   s	            r   Úcount_adap_avgpoolrC   ¦   s~   € Ü�X‰X�a˜‘d—j‘j  �nÓ%¬¯©°!·'±'¸!¸"°+Ó)>Ñ>€FÜ—‘˜“€IØ€IØ˜YÑ&€JØ—7‘7“9€LØ˜\Ñ)€IØ‡K‚K”3”s˜9“~Ó&Ñ&†Kr   c                 ó.   — | xj                   dz  c_         y r1   )r(   ©r*   r+   r,   s      r   Úcount_zero_opsrF   °   s   € Ø‡K‚K�1Ñ†Kr   c                 óš   — d}| j                  «       D ]  }||j                  «       z  }Œ t        t        |«      «      | j                  d<   y r1   )Ú
parametersr&   r)   r%   Útotal_params)r*   r+   r,   rI   Úps        r   Úcount_parametersrK   ´   s?   € Ø€LØ�\‰\Ž^ˆØ˜Ÿ™›	Ñ!‰ð äœC Ó-Ó.€A‡N�N�1Òr   c                 óX  — | j                  dt        j                  |d   j                  «      «       t	        |t
        t        f«      r3| j                  dt        j                  |d   j                  «      «       y | j                  dt        j                  |j                  «      «       y )NÚinput_shaper   Úoutput_shape)Úregister_bufferr   Ú	to_tensorr#   r   ÚlistÚtuplerE   s      r   Úcount_io_inforS   »   su   € Ø×Ñ�m¤V×%5Ñ%5°a¸±d·j±jÓ%AÔBÜ�!”dœE�]Ô#Ø	×Ñ˜.¬&×*:Ñ*:¸1¸Q¹4¿:¹:Ó*FÕGà	×Ñ˜.¬&×*:Ñ*:¸1¿7¹7Ó*CÕDr   c           
      óò  ‡‡‡— g Št        «       Š‰€i Šˆˆˆfd„}| j                  }| j                  «        | j                  |«       t        j
                  j                  «       5   | |«       d d d «       d}d}| j                  «       D ]{  }t        t        |j                  «       «      «      dkD  rŒ)h d£j                  t        |j                  j                  «       «      «      sŒ^||j                  z  }||j                  z  }Œ} |r| j!                  «        ‰D ]  }	|	j#                  «        Œ t%        g d¢«      }
| j'                  «       D �]Y  \  }}t        t        |j                  «       «      «      dkD  rŒ-h d£j                  t        |j                  j                  «       «      «      sŒb|
j)                  |j+                  «       t        |j,                  j/                  «       «      t        |j0                  j/                  «       «      t3        |j                  «      t3        |j                  «      g«       |j                  j5                  d«       |j                  j5                  d«       |j                  j5                  d«       |j                  j5                  d«       �Œ\ |r|
j7                  «        t9        d	t3        |«      › d
t3        |«      › �«       t3        |«      S # 1 sw Y   �ŒxY w)Nc                 ó´  •— t        t        | j                  «       «      «      dkD  ry | j                  dt	        j
                  dgd¬«      «       | j                  dt	        j
                  dgd¬«      «       t        | «      }d }|‰v r‰|   }|‰vrFt        d|› �«       n7|t        v rt        |   }|‰vr"t        |› d�«       n|‰vrt        d	|› d
�«       |�"| j                  |«      }‰j                  |«       | j                  t        «      }| j                  t        «      }‰j                  |«       ‰j                  |«       ‰j                  |«       y )Nr   r(   r   Úint64)ÚdtyperI   z'Customize Function has been applied to z's flops has been countedz(Cannot find suitable count function for z. Treat it as zero FLOPs.)ÚlenrQ   ÚchildrenrO   r   ÚzerosÚtypeÚprintÚregister_hooksÚregister_forward_post_hookÚappendrK   rS   Úadd)	r*   Úm_typeÚflops_fnÚflops_handlerÚparams_handlerÚ
io_handlerr
   Úhandler_collectionÚtypes_collections	         €€€r   Ú	add_hooksz dynamic_flops.<locals>.add_hooksà   sK  ø€ ÜŒt�A—J‘J“LÓ!Ó" QÒ&ØØ	×Ñ˜+¤v§|¡|°Q°C¸wÔ'GÔHØ	×Ñ˜.¬&¯,©,¸°sÀ'Ô*JÔKÜ�a“ˆàˆØ�ZÑØ! &Ñ)ˆHØÐ-Ñ-ÜÐ?À¸xÐHÕIØ”~Ñ%Ü% fÑ-ˆHØÐ-Ñ-Ü˜˜Ð 9Ð:Õ;àÐ-Ñ-ÜØ>¸v¸hÐF_Ð`ôð ÐØ×8Ñ8¸ÓBˆMØ×%Ñ% mÔ4Ø×5Ñ5Ô6FÓGˆØ×1Ñ1´-Ó@ˆ
Ø×!Ñ! .Ô1Ø×!Ñ! *Ô-Ø×Ñ˜VÕ$r   r   >   r(   rM   rN   rI   )z
Layer NamezInput ShapezOutput ShapeÚParamsÚFlopsr(   rI   rM   rN   zTotal Flops: z     Total Params: )Úsetr6   ÚevalÚapplyr   Ú	frameworkÚno_gradÚ	sublayersrX   rQ   rY   ÚissubsetÚ_buffersÚkeysr(   rI   ÚtrainÚremover   Únamed_sublayersÚadd_rowÚ	full_namerM   ÚnumpyrN   r%   ÚpopÚprint_tabler\   )Úmodelr	   r
   r   rh   r6   r(   rI   r*   ÚhandlerÚtableÚnrf   rg   s     `         @@r   r   r   Ú   sQ  ú€ ØÐÜ“uÐØÐØˆ
ö%ð> �~‰~€Hà	‡J�J„LØ	‡K�K�	Ôä	×	Ñ	×	!Ñ	!Õ	#ÙˆfŒ÷ 
$ð €IØ€LØ�_‰_ÖˆÜŒt�A—J‘J“LÓ!Ó" QÒ&Øò
÷
 ‰(”3�q—z‘z—‘Ó(Ó)Ó
*ñ+ð ˜Ÿ™Ñ$ˆIØ˜AŸN™NÑ*‰Lð ñ Ø�‰ŒÛ%ˆØ�‰Õð &ô ÚHó€Eð ×%Ñ%×'‰ˆˆ1ÜŒt�A—J‘J“LÓ!Ó" QÒ&Øò
÷
 ‰(”3�q—z‘z—‘Ó(Ó)Ó
*ñ+ð �M‰Mà—K‘K“MÜ˜Ÿ™×,Ñ,Ó.Ó/Ü˜Ÿ™×-Ñ-Ó/Ó0Ü˜Ÿ™Ó'Ü˜Ÿ™Ó$ðôð �J‰J�N‰N˜;Ô'Ø�J‰J�N‰N˜>Ô*Ø�J‰J�N‰N˜=Ô)Ø�J‰J�N‰N˜>Ö*ð+ (ñ, Ø×ÑÔÜ	Ø
œ˜I›Ð'Ð':¼3¸|Ó;LÐ:MÐNôô ˆy‹>Ð÷k 
$Ñ	#ús   Á&	K,Ë,K6)NF)+r   ry   r    r   r   Ú'paddle.jit.dy2static.program_translatorr   r   r   Ú__all__r   r/   r4   r7   r;   r=   rC   rF   rK   rS   ÚConv1DÚConv2DÚConv3DÚConv1DTransposeÚConv2DTransposeÚConv3DTransposeÚlayerÚnormÚBatchNorm2DÚ	BatchNormÚReLUÚReLU6Ú	LeakyReLUÚLinearÚDropoutÚ	AvgPool1DÚ	AvgPool2DÚ	AvgPool3DÚAdaptiveAvgPool1DÚAdaptiveAvgPool2DÚAdaptiveAvgPool3Dr]   r   © r   r   Ú<module>r˜      sŸ  ðó ã ã Ý Ý Eç -à
€ó`òF'ò"ò'ò'ò"ò'òò/òEðØ‡I�Iˆ|ðà‡I�Iˆ|ðð ‡I�Iˆ|ðð ×Ñ˜ð	ð
 ×Ñ˜ðð ×Ñ˜ðð ‡H�H‡M�M×Ñ˜xðð ‡L�L�(ðð ‡G�Gˆ^ðð ‡H�Hˆnðð ‡L�LÐ"ðð ‡I�Iˆ|ðð ‡J�J�ðð ‡L�L�-ðð ‡L�L�-ðð  ‡L�L�-ð!ð" ×ÑÐ,ð#ð$ ×ÑÐ,Ø×ÑÐ,ñ'€ô._r   