U
    hâËd  ã                   @   s:   d dl mZ d dlmZ d dlZd dlmZ ddd„ZdS )é    )Úcuda)ÚdriverN)Únumpy_supportc                    sô   t | ddƒ}|sJ| j\}}| jj| | jjf}tjjj||f|| j|d�}t 	| j¡‰ t
 ¡ j}tt dt |d¡d ¡ƒ}t|| ƒ}||d f‰tj‡ ‡fdd„ƒ}	t|jd | d ƒt|jd | d ƒf}
||f}|	|
||f | |ƒ |S )aá  Compute the transpose of 'a' and store it into 'b', if given,
    and return it. If 'b' is not given, allocate a new array
    and return that.

    This implements the algorithm documented in
    http://devblogs.nvidia.com/parallelforall/efficient-matrix-transpose-cuda-cc/

    :param a: an `np.ndarray` or a `DeviceNDArrayBase` subclass. If already on
        the device its stream will be used to perform the transpose (and to copy
        `b` to the device if necessary).
    Ústreamr   )Údtyper   é   é   c           	         sÌ   t jjˆˆ d�}t jj}t jj}t jjt jj }t jjt jj }|| }|| }|| | jd k r�|| | jd k r�| || || f |||f< t  	¡  ||jd k rÈ||jd k rÈ|||f |||f< d S )N)Úshaper   r   r   )
r   ZsharedÚarrayZ	threadIdxÚxÚyZblockIdxZblockDimr	   Zsyncthreads)	ÚinputÚoutputZtileZtxÚtyZbxZbyr   r   ©ÚdtZ
tile_shape© úU/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numba/cuda/kernels/transpose.pyÚkernel)   s    $ztranspose.<locals>.kernel)Úgetattrr	   r   Úitemsizer   ZcudadrvZdevicearrayZDeviceNDArrayÚnpsZ
from_dtyper   Z
get_deviceZMAX_THREADS_PER_BLOCKÚintÚmathÚpowÚlogZjit)ÚaÚbr   ÚcolsÚrowsÚstridesZtpbZ
tile_widthZtile_heightr   ÚblocksÚthreadsr   r   r   Ú	transpose   s*    
ü
,r#   )N)	Znumbar   Znumba.cuda.cudadrv.driverr   r   Znumba.npr   r   r#   r   r   r   r   Ú<module>   s   