o
    8ήcK"                     @   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GenerializedUFuncGUFuncCallStepsc                   @   sL   e Zd ZdZdd Zedd Zej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   || _ d| _d S )Nr   )	functions_maxblocksize)selftypes_to_retty_kernels r   =/tmp/pip-target-vg8gfxp4/lib/python/numba/cuda/vectorizers.py__init__   s   
zCUDAUFuncDispatcher.__init__c                 C   s   | j S N)r
   r   r   r   r   max_blocksize   s   z!CUDAUFuncDispatcher.max_blocksizec                 C   s
   || _ d S r   )_max_blocksize)r   blkszr   r   r   r      s   
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J d|jdksJ d|jd }g }|dkr)td|dkr1|d S |p6t }|	 0 tj
j|rF|}nt||}| |||}td|jd}|j||d	 W d    |d S 1 snw   Y  |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ndimshape	TypeErrorr   r   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r2||d \}}|| || | |||}|| | ||||dS ||d \}}	|| ||	 | ||	||d |d dkrZ| |||S |S )Nr   r   r   )r1   r   )r#   splitappendr*   )
r   r0   r/   r   r.   fatcutthincutr1   leftrightr   r   r   __reduceC   s   





zCUDAUFuncDispatcher.__reduceNr   )__name__
__module____qualname____doc__r   propertyr   setterr   r3   r*   r   r   r   r   r      s    


r   c                   @   sJ   e Zd ZdgZdd Zdd Zdd Zdd	 Zd
d Zdd Z	dd Z
dS )_CUDAGUFuncCallSteps_streamc                 C   
   t |S r   r   is_cuda_arrayr   objr   r   r   is_device_array`      
z$_CUDAGUFuncCallSteps.is_device_arrayc                 C      t jj|r	|S t |S r   r   r&   r'   r(   as_cuda_arrayrG   r   r   r   as_device_arrayc      
z$_CUDAGUFuncCallSteps.as_device_arrayc                 C   s   t j|| jdS Nr   )r   r)   rC   )r   hostaryr   r   r   r)   m      z_CUDAGUFuncCallSteps.to_devicec                 C   s   |j || jd}|S rP   )r,   rC   )r   devaryrQ   r1   r   r   r   to_hostp   s   z_CUDAGUFuncCallSteps.to_hostc                 C   s   t j||| jdS N)r#   r   r   )r   device_arrayrC   )r   r#   r   r   r   r   rV   t   s   z!_CUDAGUFuncCallSteps.device_arrayc                 C   s   | j dd| _d S )Nr   r   )kwargsgetrC   r   r   r   r   prepare_inputsw   s   z#_CUDAGUFuncCallSteps.prepare_inputsc                 C   s   |j || jd|  d S rP   )forallrC   )r   kernelnelemr   r   r   r   launch_kernelz   s   z"_CUDAGUFuncCallSteps.launch_kernelN)r<   r=   r>   	__slots__rI   rN   r)   rT   rV   rY   r]   r   r   r   r   rB   [   s    
rB   c                   @   s(   e Zd Zedd Zdd Zdd ZdS )CUDAGenerializedUFuncc                 C      t S r   )rB   r   r   r   r   _call_steps      z!CUDAGenerializedUFunc._call_stepsc                 C   s   t jjj|d|j|jdS Nr;   r#   stridesr   gpu_data)r   r&   r'   DeviceNDArrayr   rf   )r   aryr#   r   r   r   _broadcast_scalar_input   s
   
z-CUDAGenerializedUFunc._broadcast_scalar_inputc                 C   s:   t |t |j }d| |j }tjjj|||j|jdS rc   )	r   r#   re   r   r&   r'   rg   r   rf   )r   rh   newshapenewax
newstridesr   r   r   _broadcast_add_axis   s   
z)CUDAGenerializedUFunc._broadcast_add_axisN)r<   r=   r>   r@   ra   ri   rm   r   r   r   r   r_   ~   s
    
r_   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 rP   )rZ   )r   funccountr   r   r   r   r   launch   s   zCUDAUFuncMechanism.launchc                 C   rD   r   rE   rG   r   r   r   rI      rJ   z"CUDAUFuncMechanism.is_device_arrayc                 C   rK   r   rL   rG   r   r   r   rN      rO   z"CUDAUFuncMechanism.as_device_arrayc                 C   s   t j||dS rP   )r   r)   )r   rQ   r   r   r   r   r)         zCUDAUFuncMechanism.to_devicec                 C   s   |j |dS rP   )r,   )r   rS   r   r   r   r   rT      s   zCUDAUFuncMechanism.to_hostc                 C   s   t j|||dS rU   )r   rV   )r   r#   r   r   r   r   r   rV      rR   zCUDAUFuncMechanism.device_arrayc                    sn    fddt tD }tt j }dg| t j }|D ]}d||< q#tjjj| j	 j
dS )Nc                    s,   g | ]}| j ks j| | kr|qS r   )r"   r#   ).0axrh   r#   r   r   
<listcomp>   s
    
z7CUDAUFuncMechanism.broadcast_device.<locals>.<listcomp>r   rd   )ranger   r#   r    re   r   r&   r'   rg   r   rf   )r   rh   r#   
ax_differs
missingdimre   rs   r   rt   r   broadcast_device   s   

z#CUDAUFuncMechanism.broadcast_deviceN)r<   r=   r>   r?   DEFAULT_STREAMrp   rI   rN   r)   rT   rV   ry   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   jitpyfunc	overloadsr   	signaturereturn_type)r   sigcudevfnr   r   r   _compile_core   s   zCUDAVectorize._compile_corec                 C   s    | j j }|t|d |S )N__cuda____core__)r   __globals__copyupdater   )r   corefnglblr   r   r   _get_globals   s
   zCUDAVectorize._get_globalsc                 C   rD   r   r   r~   r   fnobjr   r   r   r   _compile_kernel   rJ   zCUDAVectorize._compile_kernelc                 C   s
   t | jS r   )r   	kernelmapr   r   r   r   build_ufunc   rJ   zCUDAVectorize.build_ufuncc                 C   r`   r   )vectorizer_stager_sourcer   r   r   r   _kernel_template   rb   zCUDAVectorize._kernel_templateN)	r<   r=   r>   r   r   r   r   r@   r   r   r   r   r   r{      s    r{   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|dS )N)r   engine)r   GUFuncEngineinputsig	outputsigr_   r   )r   r   r   r   r   r      s   zCUDAGUFuncVectorize.build_ufuncc                 C   s   t ||S r   r   r   r   r   r   r      rq   z#CUDAGUFuncVectorize._compile_kernelc                 C   r`   r   )_gufunc_stager_sourcer   r   r   r   r      rb   z$CUDAGUFuncVectorize._kernel_templatec                 C   s4   t j|dd| j}| jj }|t |d |S )NT)r|   r   )r   r~   r   py_funcr   r   r   )r   r   r   glblsr   r   r   r      s   z CUDAGUFuncVectorize._get_globalsN)r<   r=   r>   r   r   r@   r   r   r   r   r   r   r      s    
r   N)numbar   numpyr   r+   numba.np.ufuncr   numba.np.ufunc.deviceufuncr   r   r   objectr   rB   r_   r   r   DeviceVectorizer{   r   DeviceGUFuncVectorizer   r   r   r   r   <module>   s    S#0