Ë
    –\;j<&  ã                   ó”   — d dl Z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 ddlmZ g Ze ed	¬
«      	 	 	 	 	 	 	 	 	 dd„«       «       Zy)é    N)Ústatic_only)ÚLayerHelper)Útemplatedoc)Ú	ParamAttr)ÚAssigné   )Úcheck_variable_and_dtypeÚnce)Úop_typec                 óª  ‡#— t        d%i t        «       ¤ŽŠ#t        | dddgd«       t        |ddgd«       | j                  dk7  rt	        d| j                  › d	�«      ‚| j
                  d
   }|j
                  d
   }‰#j                  ‰#j                  ||gd| j                  ¬«      }i }‰#j                  r0‰#j                  ‰#j                  |d
gd| j                  ¬«      }||d<   ‰#j                  | j                  ¬«      }‰#j                  | j                  ¬«      }‰#j                  |j                  ¬«      }| |d<   ||d<   ||d<   |�|ng |d<   |dk(  rd}�n-|dk(  rd
}�n$|dk(  �r|	€J ‚|}dg|z  }dg|z  }g }g }t        |«      D ]L  }|	|   |z  }|dz
  dkD  r|j                  ||f«       Œ'd|z
  dkD  r|j                  ||f«       ŒC|||<   d||<   ŒN t        |«      r±t        |«      r¦|j                  d«      }|j                  d«      }|d   }|d
   }|d
   ||d   <   |||d   <   |d
   |d
   z   d
z
  }|dz
  dkD  r|j                  ||f«       n&d|z
  dkD  r|j                  ||f«       n
|||<   d||<   t        |«      rt        |«      rŒ¦t        |«      r!|j                  d«      }d||d   <   d||d   <   t        |«      r!|j                  d«      }d||d   <   d||d   <   ˆ#fd„}  | t        j                   |	«      j#                  d«      «      |d<    | t        j                   |«      j#                  d«      «      |d<    | t        j                   |«      j#                  d«      «      |d<   d}nt%        d«      ‚|€d }nt'        |«      }|}!t)        d!«       t'        |«      ||
|||!d"œ}"‰#j+                  d||||d#œ|"¬$«       ||d
z   z  S )&a�  
    :api_attr: Static Graph

    ${comment}

    Args:
        input (Tensor): Input tensor, 2-D tensor with shape [batch_size, dim],
            and data type is float32 or float64.
        label (Tensor): Input label, 2-D tensor with shape [batch_size, num_true_class],
            and data type is int64.
        num_total_classes (int):${num_total_classes_comment}.
        sample_weight (Tensor|None): A Tensor of shape [batch_size, 1]
            storing a weight for each sample. The default weight for each
            sample is 1.0.
        param_attr (ParamAttr|None): To specify the weight parameter attribute.
            Default: None, which means the default weight parameter property is
            used. See usage for details in :ref:`api_paddle_ParamAttr` .
        bias_attr (ParamAttr|None): To specify the bias parameter attribute.
            Default: None, which means the default bias parameter property is
            used. See usage for details in :ref:`api_paddle_ParamAttr` .
        num_neg_samples (int): ${num_neg_samples_comment}.
        name(str|None): For detailed information, please refer to
            :ref:`api_guide_Name` . Usually name is no need to set and None by default.
        sampler (str, optional): The sampler used to sample class from negative classes.
                       It can be 'uniform', 'log_uniform' or 'custom_dist'.
                       default: 'uniform'.
        custom_dist (nd.array|None): A numpy ndarray with size=num_total_classes.
                       It is used when sampler is set to 'custom_dist'.
                       custom_dist[i] is the probability of i-th class to be sampled.
                       default: None.
        seed (int, optional): The seed used in sampler. Default 0, means no random seed.
        is_sparse(bool, optional): The flag indicating whether to use sparse update,
            the weight@GRAD and bias@GRAD will be changed to SelectedRows. Default False.

    Returns:
        Tensor: The output nce loss.

    Examples:
        .. code-block:: python

            >>> import paddle
            >>> import numpy as np

            >>> paddle.enable_static()

            >>> window_size = 5
            >>> words = []
            >>> for i in range(window_size):
            ...     words.append(paddle.static.data(
            ...         name='word_{0}'.format(i), shape=[-1, 1], dtype='int64'))

            >>> dict_size = 10000
            >>> label_word = int(window_size / 2) + 1

            >>> embs = []
            >>> for i in range(window_size):
            ...     if i == label_word:
            ...         continue
            ...
            ...     emb = paddle.static.nn.embedding(input=words[i], size=[dict_size, 32],
            ...                         param_attr='embed', is_sparse=True)
            ...     embs.append(emb)

            >>> embs = paddle.concat(x=embs, axis=1)                # concat from 4 * [(-1, 1, 32)] to (-1, 4, 32)
            >>> embs = paddle.reshape(x=embs, shape=(-1, 4 * 32))   # reshape to (batch_size = -1, dim = 4*32)
            >>> loss = paddle.static.nn.nce(input=embs, label=words[label_word],
            ...             num_total_classes=dict_size, param_attr='nce.w_0',
            ...             bias_attr='nce.b_0')

            # or use custom distribution
            >>> dist = np.array([0.05,0.5,0.1,0.3,0.05])
            >>> loss = paddle.static.nn.nce(input=embs, label=words[label_word],
            ...         num_total_classes=5, param_attr='nce.w_1',
            ...         bias_attr='nce.b_1',
            ...         num_neg_samples=3,
            ...         sampler="custom_dist",
            ...         custom_dist=dist)
    r
   ÚinputÚfloat32Úfloat64ÚlabelÚint64é   z,The rank of `input` must be 2, but received Ú.é   F)ÚattrÚshapeÚis_biasÚdtypeTÚBias)r   ÚInputÚLabelÚWeightÚSampleWeightÚuniformr   Úlog_uniformÚcustom_distg      ð?éÿÿÿÿc                 óŠ   •— ‰j                  t        «       | j                  | j                  t	        | «      ¬«      }d|_        |S )N)r   r   r   Údefault_initializerT)Úcreate_parameterr   r   r   r   Ústop_gradient)Únumpy_arrayÚretÚhelpers     €ú^G:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/static/nn/loss.pyÚ_init_by_numpy_arrayz!nce.<locals>._init_by_numpy_arrayÓ   sF   ø€ Ø×)Ñ)Ü“[Ø!×'Ñ'Ø!×'Ñ'Ü$*¨;Ó$7ð	 *ó ˆCð !%ˆCÔØˆJó    ÚCustomDistProbsÚint32ÚCustomDistAliasÚCustomDistAliasProbszUnsupported sampler type.é
   zWWith sparse mode, if your models has only small parameter prefetch may cause speed down)Únum_total_classesÚnum_neg_samplesÚseedÚsamplerÚ	is_sparseÚremote_prefetch)ÚCostÚSampleLogitsÚSampleLabels)ÚtypeÚinputsÚoutputsÚattrs)r
   )r   Úlocalsr	   ÚndimÚ
ValueErrorr   r$   Ú
param_attrr   Ú	bias_attrÚ"create_variable_for_type_inferenceÚrangeÚappendÚlenÚpopÚnpÚarrayÚastypeÚ	ExceptionÚintÚprintÚ	append_op)$r   r   r1   Úsample_weightrA   rB   r2   Únamer4   r    r3   r5   ÚdimÚnum_true_classÚwr;   ÚbÚcostÚsample_logitsÚsample_labelsÚcustom_dist_lenÚalias_probs_Úalias_ÚbigsÚlittlesÚiÚnormal_probÚbigÚlittleÚbig_idxÚbig_probÚbig_leftr*   r6   r=   r(   s$                                      @r)   r
   r
   !   sq  ø€ ô| Ñ+¤&£(Ñ+€FÜ˜U G¨i¸Ð-CÀUÔKÜ˜U G¨g¨Y¸Ô>à‡z�z�Q‚ÜØ:¸5¿:¹:¸,ÀaÐHó
ð 	
ð �+‰+�a‰.€CØ—[‘[ ‘^€NØ×ÑØ×ÑØ  #Ð&ØØ�k‰kð	 	 ó 	€Að €FØ×ÒØ×#Ñ#Ø×!Ñ!Ø$ aÐ(ØØ—+‘+ð	 $ó 
ˆð ˆˆv‰Ø×4Ñ4¸5¿;¹;Ð4ÓG€DØ×=Ñ=ÀEÇKÁKÐ=ÓP€MØ×=Ñ=ÀEÇKÁKÐ=ÓP€Mà€Fˆ7�OØ€Fˆ7�OØ€Fˆ8ÑØ.;Ð.G™]ÈR€Fˆ>Ñà�)ÒØŠØ	�MÒ	!ØŠØ	�MÓ	!ØÐ&Ð&Ð&à+ˆØ�s˜_Ñ,ˆØ��Ñ&ˆØˆØˆÜ�Ö'ˆAØ% a™.¨?Ñ:ˆKØ˜SÑ  1Ò$Ø—‘˜Q Ð,Õ-Ø�{Ñ" QÒ&Ø—‘  ;Ð/Õ0à"-�˜Q‘Ø��q’	ð (ô �$ŒiœC œLØ—(‘(˜1“+ˆCØ—[‘[ “^ˆFà˜!‘fˆGØ˜1‘vˆHà&,¨Q¡iˆL˜ ™Ñ#Ø 'ˆF�6˜!‘9ÑØ˜1‘v  q¡	Ñ)¨AÑ-ˆHØ˜#‰~ Ò!Ø—‘˜W hÐ/Õ0Ø�x‘ !Ò#Ø—‘ ¨Ð2Õ3à(0�˜WÑ%Ø"$��w‘ô! �$ŒiœC �Lô$ ˆtŒ9Ø—(‘(˜1“+ˆCØ#&ˆL˜˜Q™Ñ ØˆF�3�q‘6‰NÜˆwŒ<Ø—[‘[ “^ˆFØ&)ˆL˜ ™Ñ#Ø "ˆF�6˜!‘9Ñô	ñ %9Ü�H‰H�[Ó!×(Ñ(¨Ó3ó%
ˆÐ Ñ!ñ %9Ü�H‰H�VÓ×#Ñ# GÓ,ó%
ˆÐ Ñ!ñ *>Ü�H‰H�\Ó"×)Ñ)¨)Ó4ó*
ˆÐ%Ñ&ð ‰äÐ3Ó4Ð4àÐØ‰ä˜oÓ.ˆà€OÜ	Øaôô
 !Ð!2Ó3Ø*ØØØØ*ñ€Eð ×ÑØØàØ)Ø)ñ
ð
 ð ô 	ð �? QÑ&Ñ'Ð'r+   )	NNNNNr   Nr   F)ÚnumpyrH   Úpaddle.base.frameworkr   Úpaddle.base.layer_helperr   Ú+paddle.base.layers.layer_function_generatorr   Úpaddle.base.param_attrr   Úpaddle.nn.initializerr   Úbase.data_feederr	   Ú__all__r
   © r+   r)   Ú<module>rm      sc   ðó å -õ 1Ý CÝ ,Ý (å 8à
€ð Ù�UÔð
 ØØØØ	ØØØ	
Øòd(ó ó ñd(r+   