U
    hâËdß"  ã                   @   s¨   d dl mZ d dlmZ d dlmZ d dlmZm	Z	m
Z
 G dd„ deƒZG dd„ de
ƒZG d	d
„ d
e	ƒZG dd„ deƒZdZG dd„ dejƒZdZG dd„ dejƒZdS )é    )Úcuda)Úarray)Údeviceufunc)ÚUFuncMechanismÚGeneralizedUFuncÚGUFuncCallStepsc                   @   s2   e Zd ZdZdd„ Zdd„ Zddd„Zd	d
„ ZdS )ÚCUDAUFuncDispatcherzD
    Invoke the CUDA ufunc specialization for the given inputs.
    c                 C   s   || _ |j| _d S ©N)Ú	functionsÚ__name__)ÚselfZtypes_to_retty_kernelsÚpyfunc© r   úO/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numba/cuda/vectorizers.pyÚ__init__   s    zCUDAUFuncDispatcher.__init__c                 O   s   t  | j||¡S )a¦  
        *args: numpy arrays or DeviceArrayBase (created by cuda.to_device).
               Cannot mix the two types in one call.

        **kws:
            stream -- cuda stream; when defined, asynchronous mode is used.
            out    -- output array. Can be a numpy array or DeviceArrayBase
                      depending on the input arguments.  Type must match
                      the input arguments.
        )ÚCUDAUFuncMechanismÚcallr
   )r   ÚargsÚkwsr   r   r   Ú__call__   s    zCUDAUFuncDispatcher.__call__r   c              	   C   sÖ   t t| j ¡ ƒd ƒdks"tdƒ‚|jdks4tdƒ‚|jd }g }|dkrTtdƒ‚n|dkrd|d S |pnt 	¡ }| 
¡ �P tjj |¡rŽ|}nt ||¡}|  |||¡}td|jd�}|j||d	� W 5 Q R X |d S )
Nr   é   zmust be a binary ufuncé   zmust use 1d arrayzReduction on an empty array.)r   )Údtype©Ústream)ÚlenÚlistr
   ÚkeysÚAssertionErrorÚndimÚshapeÚ	TypeErrorr   r   Zauto_synchronizeÚcudadrvÚdevicearrayÚis_cuda_ndarrayÚ	to_deviceÚ_CUDAUFuncDispatcher__reduceÚnp_arrayr   Úcopy_to_host)r   Úargr   ÚnÚgpu_memsÚmemÚoutÚbufr   r   r   Úreduce   s"    "


zCUDAUFuncDispatcher.reducec           
      C   s¼   |j d }|d dkrd| |d ¡\}}| |¡ | |¡ |  |||¡}| |¡ | ||||d�S | |d ¡\}}	| |¡ | |	¡ | ||	||d� |d dkr´|  |||¡S |S d S )Nr   r   r   )r-   r   )r    ÚsplitÚappendr&   )
r   r,   r+   r   r*   ZfatcutZthincutr-   ÚleftÚrightr   r   r   Z__reduce;   s    





zCUDAUFuncDispatcher.__reduceN)r   )r   Ú
__module__Ú__qualname__Ú__doc__r   r   r/   r&   r   r   r   r   r      s
   
r   c                       sR   e Zd ZdgZ‡ fdd„Zdd„ Zdd„ Zdd	„ Zd
d„ Zdd„ Z	dd„ Z
‡  ZS )Ú_CUDAGUFuncCallStepsÚ_streamc                    s$   t ƒ  ||||¡ | dd¡| _d S )Nr   r   )Úsuperr   Úgetr8   )r   ZninZnoutr   Úkwargs©Ú	__class__r   r   r   X   s    z_CUDAGUFuncCallSteps.__init__c                 C   s
   t  |¡S r	   ©r   Zis_cuda_array©r   Úobjr   r   r   Úis_device_array\   s    z$_CUDAGUFuncCallSteps.is_device_arrayc                 C   s   t jj |¡r|S t  |¡S r	   ©r   r"   r#   r$   Zas_cuda_arrayr?   r   r   r   Úas_device_array_   s    z$_CUDAGUFuncCallSteps.as_device_arrayc                 C   s   t j|| jd�S ©Nr   )r   r%   r8   )r   Úhostaryr   r   r   r%   i   s    z_CUDAGUFuncCallSteps.to_devicec                 C   s   |j || jd�}|S rD   )r(   r8   )r   ÚdevaryrE   r-   r   r   r   Úto_hostl   s    z_CUDAGUFuncCallSteps.to_hostc                 C   s   t j||| jd�S ©N)r    r   r   )r   Údevice_arrayr8   )r   r    r   r   r   r   Úallocate_device_arrayp   s    z*_CUDAGUFuncCallSteps.allocate_device_arrayc                 C   s   |j || jd�|Ž  d S rD   )Úforallr8   )r   ZkernelZnelemr   r   r   r   Úlaunch_kernels   s    z"_CUDAGUFuncCallSteps.launch_kernel)r   r4   r5   Ú	__slots__r   rA   rC   r%   rG   rJ   rL   Ú__classcell__r   r   r<   r   r7   S   s   ÿ
r7   c                       s8   e Zd Z‡ fdd„Zedd„ ƒZdd„ Zdd„ Z‡  ZS )	ÚCUDAGeneralizedUFuncc                    s   |j | _ tƒ  ||¡ d S r	   )r   r9   r   )r   Ú	kernelmapÚenginer   r<   r   r   r   x   s    zCUDAGeneralizedUFunc.__init__c                 C   s   t S r	   )r7   ©r   r   r   r   Ú_call_steps|   s    z CUDAGeneralizedUFunc._call_stepsc                 C   s   t jjj|d|j|jd�S ©N)r   ©r    Ústridesr   Úgpu_data)r   r"   r#   ÚDeviceNDArrayr   rW   )r   Úaryr    r   r   r   Ú_broadcast_scalar_input€   s
    
ýz,CUDAGeneralizedUFunc._broadcast_scalar_inputc                 C   s:   t |ƒt |jƒ }d| |j }tjjj|||j|jd�S rT   )	r   r    rV   r   r"   r#   rX   r   rW   )r   rY   ZnewshapeZnewaxZ
newstridesr   r   r   Ú_broadcast_add_axis†   s    
ýz(CUDAGeneralizedUFunc._broadcast_add_axis)	r   r4   r5   r   ÚpropertyrS   rZ   r[   rN   r   r   r<   r   rO   w   s
   
rO   c                   @   sL   e Zd Zd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S )r   z%
    Provide CUDA specialization
    r   c                 C   s   |j ||d�|Ž  d S rD   )rK   )r   ÚfuncÚcountr   r   r   r   r   Úlaunch–   s    zCUDAUFuncMechanism.launchc                 C   s
   t  |¡S r	   r>   r?   r   r   r   rA   ™   s    z"CUDAUFuncMechanism.is_device_arrayc                 C   s   t jj |¡r|S t  |¡S r	   rB   r?   r   r   r   rC   œ   s    z"CUDAUFuncMechanism.as_device_arrayc                 C   s   t j||d�S rD   )r   r%   )r   rE   r   r   r   r   r%   ¦   s    zCUDAUFuncMechanism.to_devicec                 C   s   |j |d�S rD   )r(   )r   rF   r   r   r   r   rG   ©   s    zCUDAUFuncMechanism.to_hostc                 C   s   t j|||d�S rH   )r   rI   )r   r    r   r   r   r   r   rJ   ¬   s    z(CUDAUFuncMechanism.allocate_device_arrayc                    sn   ‡ ‡fdd„t tˆƒƒD ƒ}tˆƒtˆ jƒ }dg| tˆ jƒ }|D ]}d||< qFtjjjˆ|ˆ j	ˆ j
d�S )Nc                    s,   g | ]$}|ˆ j ks$ˆ j| ˆ| kr|‘qS r   )r   r    )Ú.0Úax©rY   r    r   r   Ú
<listcomp>°   s    
þz7CUDAUFuncMechanism.broadcast_device.<locals>.<listcomp>r   rU   )Úranger   r    r   rV   r   r"   r#   rX   r   rW   )r   rY   r    Z
ax_differsZ
missingdimrV   ra   r   rb   r   Úbroadcast_device¯   s    

ýz#CUDAUFuncMechanism.broadcast_deviceN)r   r4   r5   r6   ZDEFAULT_STREAMr_   rA   rC   r%   rG   rJ   re   r   r   r   r   r   �   s   
r   z�
def __vectorized_{name}({args}, __out__):
    __tid__ = __cuda__.grid(1)
    if __tid__ < __out__.shape[0]:
        __out__[__tid__] = __core__({argitems})
c                   @   s8   e Zd Zdd„ Zdd„ Zdd„ Zdd„ Zed	d
„ ƒZdS )ÚCUDAVectorizec                 C   s*   t j|ddd�| jƒ}||j|j jjfS )NT)ÚdeviceÚinline)r   Újitr   Z	overloadsr   Ú	signatureÚreturn_type)r   ÚsigZcudevfnr   r   r   Ú_compile_coreÉ   s    zCUDAVectorize._compile_corec                 C   s    | j j ¡ }| t|dœ¡ |S )N©Z__cuda__Z__core__)r   Ú__globals__ÚcopyÚupdater   )r   ÚcorefnZglblr   r   r   Ú_get_globalsÍ   s
    ÿzCUDAVectorize._get_globalsc                 C   s
   t  |¡S r	   ©r   ri   ©r   Zfnobjrl   r   r   r   Ú_compile_kernelÓ   s    zCUDAVectorize._compile_kernelc                 C   s   t | j| jƒS r	   )r   rP   r   rR   r   r   r   Úbuild_ufuncÖ   s    zCUDAVectorize.build_ufuncc                 C   s   t S r	   )Úvectorizer_stager_sourcerR   r   r   r   Ú_kernel_templateÙ   s    zCUDAVectorize._kernel_templateN)	r   r4   r5   rm   rs   rv   rw   r\   ry   r   r   r   r   rf   È   s   rf   zy
def __gufunc_{name}({args}):
    __tid__ = __cuda__.grid(1)
    if __tid__ < {checkedarg}:
        __core__({argitems})
c                   @   s0   e Zd Zdd„ Zdd„ Zedd„ ƒZdd„ Zd	S )
ÚCUDAGUFuncVectorizec                 C   s"   t  | j| j¡}t| j|| jd�S )N)rP   rQ   r   )r   ZGUFuncEngineZinputsigZ	outputsigrO   rP   r   )r   rQ   r   r   r   rw   ê   s
    þzCUDAGUFuncVectorize.build_ufuncc                 C   s   t  |¡|ƒS r	   rt   ru   r   r   r   rv   ð   s    z#CUDAGUFuncVectorize._compile_kernelc                 C   s   t S r	   )Ú_gufunc_stager_sourcerR   r   r   r   ry   ó   s    z$CUDAGUFuncVectorize._kernel_templatec                 C   s4   t j|dd�| jƒ}| jj ¡ }| t |dœ¡ |S )NT)rg   rn   )r   ri   r   Zpy_funcro   rp   rq   )r   rl   rr   Zglblsr   r   r   rs   ÷   s    ÿz CUDAGUFuncVectorize._get_globalsN)r   r4   r5   rw   rv   r\   ry   rs   r   r   r   r   rz   é   s
   
rz   N)Znumbar   Únumpyr   r'   Znumba.np.ufuncr   Znumba.np.ufunc.deviceufuncr   r   r   Úobjectr   r7   rO   r   rx   ZDeviceVectorizerf   r{   ZDeviceGUFuncVectorizerz   r   r   r   r   Ú<module>   s   K$0