U
    hâËda)  ã                   @   sØ   d dl mZ d dlZd dlZd dlZd dlZddlmZm	Z	 ddl
mZmZmZ ddlmZ ddlmZmZ daed	d
„ ƒZdd„ ZG dd„ dƒZG dd„ deƒZG dd„ deƒZG dd„ dejƒZG dd„ deƒZdS )é    )ÚcontextmanagerNé   )ÚFakeCUDAArrayÚFakeWithinKernelCUDAArray)ÚDim3ÚFakeCUDAModuleÚswapped_cuda_moduleé   )Únormalize_kernel_dimensions)Úwrap_argÚArgHintc                 c   s*   t dkstdƒ‚| a z
dV  W 5 da X dS )z*
    Push the current kernel context.
    Nz)concurrent simulated kernel not supported)Ú_kernel_contextÚAssertionError)Úmod© r   úT/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numba/cuda/simulator/kernel.pyÚ_push_kernel_context   s
    
r   c                   C   s   t S )zT
    Get the current kernel context. This is usually done by a device function.
    )r   r   r   r   r   Ú_get_kernel_context$   s    r   c                   @   s   e Zd ZdZdd„ ZdS )ÚFakeOverloadzE
    Used only to provide the max_cooperative_grid_blocks method
    c                 C   s   dS )Nr   r   )ÚselfZblockdimr   r   r   Úmax_cooperative_grid_blocks/   s    z(FakeOverload.max_cooperative_grid_blocksN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r   r   r   +   s   r   c                   @   s   e Zd Zdd„ ZdS )ÚFakeOverloadDictc                 C   s   t ƒ S ©N)r   )r   Úkeyr   r   r   Ú__getitem__6   s    zFakeOverloadDict.__getitem__N)r   r   r   r   r   r   r   r   r   5   s   r   c                   @   sb   e Zd ZdZdg dfdd„Zdd„ Zdd„ Zd	d
„ Zdd„ Zddd„Z	e
dd„ ƒZe
dd„ ƒZdS )ÚFakeCUDAKernelz(
    Wraps a @cuda.jit-ed function.
    Fc                 C   sJ   || _ || _|| _|| _t|ƒ| _d | _d | _d| _d| _	t
 | |¡ d S )Nr   )ÚfnÚ_deviceZ	_fastmathÚ_debugÚlistÚ
extensionsÚgrid_dimÚ	block_dimÚstreamÚdynshared_sizeÚ	functoolsÚupdate_wrapper)r   r    ZdeviceZfastmathr$   Údebugr   r   r   Ú__init__A   s    
zFakeCUDAKernel.__init__c           	   
      sè   ˆj r2tˆjtƒ ƒ� ˆj|Ž W  5 Q R £ S Q R X tˆjˆjƒ\}}t||ˆjƒ}t	|ƒ�„ g ‰‡‡fdd„‰ ‡ fdd„|D ƒ}tˆj|ƒ�8 t
j|Ž D ]&}tˆj||ˆjƒ}|j|f|žŽ  q˜W 5 Q R X ˆD ]
}|ƒ  qÎW 5 Q R X d S )Nc                    sŒ   t  ‡ fdd„ˆjd | f¡\}} t| tjƒrF| jdkrFt| ƒ ˆ ¡}n0t| t	ƒr\|  ˆ ¡}nt| tj
ƒrrt| ƒ}n| }t|tƒrˆt|ƒS |S )Nc                    s   |j | dˆ dœŽS )Nr   )r'   Úretr)Zprepare_args)Zty_valÚ	extension)r-   r   r   Ú<lambda>b   s   ýz;FakeCUDAKernel.__call__.<locals>.fake_arg.<locals>.<lambda>r   )r)   Úreducer$   Ú
isinstanceÚnpZndarrayÚndimr   Z	to_devicer   Úvoidr   r   )ÚargÚ_Úret)r-   r   r   r   Úfake_arg_   s    
ú	


z)FakeCUDAKernel.__call__.<locals>.fake_argc                    s   g | ]}ˆ |ƒ‘qS r   r   )Ú.0r5   )r8   r   r   Ú
<listcomp>v   s     z+FakeCUDAKernel.__call__.<locals>.<listcomp>)r!   r   r    r   r
   r%   r&   r   r(   r   r2   ÚndindexÚBlockManagerr"   Úrun)	r   Úargsr%   r&   Zfake_cuda_moduleZ	fake_argsÚ
grid_pointZbmÚwbr   )r8   r-   r   r   Ú__call__O   s&    ÿÿ
zFakeCUDAKernel.__call__c                 C   s2   t |d d… Ž \| _| _t|ƒdkr.|d | _| S )Nr	   é   é   )r
   r%   r&   Úlenr(   )r   Úconfigurationr   r   r   r   €   s
    ÿ

zFakeCUDAKernel.__getitem__c                 C   s   d S r   r   ©r   r   r   r   Úbind‰   s    zFakeCUDAKernel.bindc                 G   s   | S r   r   )r   r>   r   r   r   Ú
specializeŒ   s    zFakeCUDAKernel.specializer   c                 C   s$   |dk rt d| ƒ‚| |d||f S )Nr   z0Can't create ForAll with negative task count: %sr   )Ú
ValueError)r   ZntasksZtpbr'   Z	sharedmemr   r   r   Úforall�   s
    ÿzFakeCUDAKernel.forallc                 C   s   t ƒ S r   )r   rF   r   r   r   Ú	overloads•   s    zFakeCUDAKernel.overloadsc                 C   s   | j S r   )r    rF   r   r   r   Úpy_func™   s    zFakeCUDAKernel.py_funcN)r   r   r   )r   r   r   r   r,   rA   r   rG   rH   rJ   ÚpropertyrK   rL   r   r   r   r   r   <   s   1	

r   c                       sT   e Zd ZdZ‡ fdd„Z‡ fdd„Zdd„ Zdd	„ Zd
d„ Zdd„ Z	dd„ Z
‡  ZS )ÚBlockThreadzG
    Manages the execution of a function for a single CUDA thread.
    c           	         s¤   |r‡ fdd„}|}nˆ }t t| ƒj|d� t ¡ | _d| _|| _t|Ž | _	t|Ž | _
d | _d| _d| _|| _t| jjŽ }| j
j|j| j
j|j| j
j    | _d S )Nc                     s   t jdd� ˆ | |Ž d S )NÚraise)Údivide)r2   Zseterr)r>   Úkwargs©Úfr   r   Údebug_wrapper¦   s    z+BlockThread.__init__.<locals>.debug_wrapper)ÚtargetFT)ÚsuperrN   r,   Ú	threadingÚEventÚsyncthreads_eventÚsyncthreads_blockedÚ_managerr   ÚblockIdxÚ	threadIdxÚ	exceptionÚdaemonÚabortr+   Ú
_block_dimÚxÚyÚzÚ	thread_id)	r   rS   Úmanagerr\   r]   r+   rT   rU   ZblockDim©Ú	__class__rR   r   r,   ¤   s(    


ÿÿzBlockThread.__init__c              
      sœ   zt t| ƒ ¡  W n„ tk
r– } zfdt| jƒ }dt| jƒ }t|ƒdkrZd||f }nd|||f }t 	¡ d }t
|ƒ|ƒ|f| _W 5 d }~X Y nX d S )Nztid=%szctaid=%sÚ z%s %sz	%s %s: %sr	   )rV   rN   r=   Ú	Exceptionr#   r]   r\   ÚstrÚsysÚexc_infoÚtyper^   )r   ÚeÚtidZctaidÚmsgÚtbrg   r   r   r=   ¼   s    zBlockThread.runc                 C   s:   | j rtdƒ‚d| _| j ¡  | j ¡  | j r6tdƒ‚d S )Nz"abort flag set on syncthreads callTz#abort flag set on syncthreads clear)r`   ÚRuntimeErrorrZ   rY   ÚwaitÚclearrF   r   r   r   ÚsyncthreadsË   s    

zBlockThread.syncthreadsc                 C   sD   | j j| j j| j jf}|| jj|< |  ¡  t | jj¡}|  ¡  |S r   )	r]   rb   rc   rd   r[   Úblock_staterv   r2   Zcount_nonzero)r   ÚvalueÚidxÚcountr   r   r   Úsyncthreads_count×   s    zBlockThread.syncthreads_countc                 C   sL   | j j| j j| j jf}|| jj|< |  ¡  t | jj¡}|  ¡  |rHdS dS ©Nr   r   )	r]   rb   rc   rd   r[   rw   rv   r2   Úall©r   rx   ry   Útestr   r   r   Úsyncthreads_andß   s    zBlockThread.syncthreads_andc                 C   sL   | j j| j j| j jf}|| jj|< |  ¡  t | jj¡}|  ¡  |rHdS dS r|   )	r]   rb   rc   rd   r[   rw   rv   r2   Úanyr~   r   r   r   Úsyncthreads_orç   s    zBlockThread.syncthreads_orc                 C   s   d| j | jf S )NzThread <<<%s, %s>>>)r\   r]   rF   r   r   r   Ú__str__ï   s    zBlockThread.__str__)r   r   r   r   r,   r=   rv   r{   r€   r‚   rƒ   Ú__classcell__r   r   rg   r   rN       s   rN   c                   @   s    e Zd ZdZdd„ Zdd„ ZdS )r<   aç  
    Manages the execution of a thread block.

    When run() is called, all threads are started. Each thread executes until it
    hits syncthreads(), at which point it sets its own syncthreads_blocked to
    True so that the BlockManager knows it is blocked. It then waits on its
    syncthreads_event.

    The BlockManager polls threads to determine if they are blocked in
    syncthreads(). If it finds a blocked thread, it adds it to the set of
    blocked threads. When all threads are blocked, it unblocks all the threads.
    The thread are unblocked by setting their syncthreads_blocked back to False
    and setting their syncthreads_event.

    The polling continues until no threads are alive, when execution is
    complete.
    c                 C   s.   || _ || _|| _|| _tj|tjd�| _d S )N)Zdtype)Z	_grid_dimra   Ú_fr"   r2   ZzerosZbool_rw   )r   rS   r%   r&   r+   r   r   r   r,     s
    zBlockManager.__init__c           
         s"  t ƒ }t ƒ }t ƒ }tjˆjŽ D ]@}‡ ‡fdd„}t|ˆ||ˆjƒ}| ¡  | |¡ | |¡ q|rø|D ]R}|jr~| |¡ qh|j	rh|D ]}	d|	_
d|	_|	j  ¡  qˆ|j	d  |j	d ¡‚qh||krä|D ]}d|_|j  ¡  qÈt ƒ }t dd„ |D ƒƒ}q`|D ] }|j	rü|j	d  |j	d ¡‚qüd S )	Nc                      s   ˆj ˆ Ž  d S r   )r…   r   ©r>   r   r   r   rU     s    z BlockManager.run.<locals>.targetTFr   r   c                 S   s   g | ]}|  ¡ r|‘qS r   )Úis_alive)r9   Útr   r   r   r:   /  s      z$BlockManager.run.<locals>.<listcomp>)Úsetr2   r;   ra   rN   r"   ÚstartÚaddrZ   r^   r`   rY   Úwith_traceback)
r   r?   r>   ÚthreadsZlivethreadsZblockedthreadsZblock_pointrU   rˆ   Zt_otherr   r†   r   r=     s8    
zBlockManager.runN)r   r   r   r   r,   r=   r   r   r   r   r<   ó   s   r<   )Ú
contextlibr   r)   rl   rW   Únumpyr2   Zcudadrv.devicearrayr   r   Z	kernelapir   r   r   Úerrorsr
   r>   r   r   r   r   r   r   Údictr   Úobjectr   ÚThreadrN   r<   r   r   r   r   Ú<module>   s"   

dS