U
    hâËd-x  ã                   @   s2  d Z ddlmZmZ ddlmZ ddlZddlZddlm	Z	 ddl
ZddlmZmZ ddlmZmZ ddlmZ dd	lmZ d
d„ Zdd„ Zdd„ ZG dd„ deƒZdd„ ZG dd„ deƒZG dd„ deƒZdd„ Zdd„ Zdd„ Z dd„ Z!G d d!„ d!eƒZ"G d"d#„ d#eƒZ#G d$d%„ d%eƒZ$G d&d'„ d'ed(�Z%dS ))zA
Implements custom ufunc dispatch mechanism for non-CPU devices.
é    )ÚABCMetaÚabstractmethod)ÚOrderedDictN)Úreduce)Ú_BaseUFuncBuilderÚparse_identity)ÚtypesÚsigutils)Ú	signature©Úparse_signaturec                 C   s8   | |kr| S | dkr|S |dkr$| S t d | |¡ƒ‚dS )ú=
    Raises
    ------
    ValueError if broadcast fails
    é   zfailed to broadcast {0} and {1}N)Ú
ValueErrorÚformat)ÚaÚb© r   úS/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numba/np/ufunc/deviceufunc.pyÚ_broadcast_axis   s    r   c                 C   s^   t t| |gƒ\} }t| ƒt|ƒk r,d|  } qt| ƒt|ƒkrFd| }q,tdd„ t| |ƒD ƒƒS )r   ©r   c                 s   s   | ]\}}t ||ƒV  qd S ©N)r   )Ú.0r   r   r   r   r   Ú	<genexpr>1   s     z&_pairwise_broadcast.<locals>.<genexpr>)ÚmapÚtupleÚlenÚzip)Zshape1Zshape2r   r   r   Ú_pairwise_broadcast#   s    

r   c                  G   sl   | st ‚| d }| dd… }z$t|dd�D ]\}}t||ƒ}q*W n" tk
rb   td |¡ƒ‚Y nX |S dS )r   r   r   N)Ústartz!failed to broadcast argument #{0})ÚAssertionErrorÚ	enumerater   r   r   )Ú	shapelistÚresultZothersÚiZeachr   r   r   Ú_multi_broadcast4   s    r%   c                   @   s¤   e Zd 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d„ Zdd„ Zdd„ Zdd„ Zdd„ Zedd„ ƒZdd„ Zd d!„ Zd"d#„ Zd$d%„ ZdS )&ÚUFuncMechanismz0
    Prepare ufunc arguments for vectorize.
    NFc                 C   s>   || _ || _t| jƒ}dg| | _g | _d| _dg| | _dS )zFNever used directly by user. Invoke by UFuncMechanism.call().
        N)ÚtypemapÚargsr   ÚargtypesÚ	scalarposr
   Úarrays)Úselfr'   r(   Únargsr   r   r   Ú__init__N   s    
zUFuncMechanism.__init__c                 C   sf   t | jƒD ]V\}}|  |¡r.|  |¡| j|< q
t|tttt	j
fƒrP| j |¡ q
t	 |¡| j|< q
dS )z1
        Get all arguments in array form
        N)r!   r(   Úis_device_arrayÚas_device_arrayr+   Ú
isinstanceÚintÚfloatÚcomplexÚnpÚnumberr*   ÚappendÚasarray)r,   r$   Úargr   r   r   Ú_fill_arraysY   s    
zUFuncMechanism._fill_arraysc                 C   sH   t | jƒD ]8\}}|dk	r
t|dƒ}|dkr8t |¡j}|| j|< q
dS )z
        Get dtypes
        NÚdtype)r!   r+   Úgetattrr5   r8   r;   r)   )r,   r$   Úaryr;   r   r   r   Ú_fill_argtypesf   s    
zUFuncMechanism._fill_argtypesc                 C   sÜ   g }| j rr| jD ]`}g }tt|| jƒƒD ]4\}\}}|dkrNt | j| ¡j}| 	||k¡ q(t
|ƒr| 	|¡ q|s®g }| jD ],}t
dd„ t|| jƒD ƒƒ}|r€| 	|¡ q€|sºtdƒ‚t|ƒdkrÎtdƒ‚|d | _dS )z<Resolve signature.
        May have ambiguous case.
        Nc                 s   s"   | ]\}}|d kp||kV  qd S r   r   )r   ÚformalÚactualr   r   r   r   ‰   s   ÿz4UFuncMechanism._resolve_signature.<locals>.<genexpr>z…No matching version.  GPU ufunc requires array arguments to have the exact types.  This behaves like regular ufunc with casting='no'.r   zqFailed to resolve ufunc due to ambiguous signature. Too many untyped scalars. Use numpy dtype object to type tag.r   )r*   r'   r!   r   r)   r5   r8   r(   r;   r7   ÚallÚ	TypeErrorr   )r,   ÚmatchesZ	formaltysZ	match_mapr$   r?   r@   Zall_matchesr   r   r   Ú_resolve_signatureq   s2    
ÿ

þz!UFuncMechanism._resolve_signaturec                 C   s4   | j D ]&}tj| j| g| j| d�| j|< q| jS )zPReturn the actual arguments
        Casts scalar arguments to np.array.
        ©r;   )r*   r5   Úarrayr(   r)   r+   )r,   r$   r   r   r   Ú_get_actual_argsœ   s    
$zUFuncMechanism._get_actual_argsc           	         sÊ   dd„ |D ƒ}t |Ž ‰t|ƒD ]¦\}‰ ˆ jˆkr2q|  ˆ ¡rN|  ˆ ˆ¡||< q‡ ‡fdd„ttˆƒƒD ƒ}tˆƒtˆ jƒ }dg| tˆ jƒ }|D ]}d||< q”t	j
jjˆ ˆ|d�}|  |¡||< q|S )z)Perform numpy ufunc broadcasting
        c                 S   s   g | ]
}|j ‘qS r   ©Úshape©r   r   r   r   r   Ú
<listcomp>¨   s     z-UFuncMechanism._broadcast.<locals>.<listcomp>c                    s,   g | ]$}|ˆ j ks$ˆ j| ˆ| kr|‘qS r   )ÚndimrI   )r   Úax©r=   rI   r   r   rK   ´   s    
þr   )rI   Ústrides)r%   r!   rI   r/   Úbroadcast_deviceÚranger   ÚlistrO   r5   ÚlibZstride_tricksZ
as_stridedÚforce_array_layout)	r,   Úarysr"   r$   Z
ax_differsZ
missingdimrO   rM   Zstridedr   rN   r   Ú
_broadcast¥   s$    



þzUFuncMechanism._broadcastc                 C   s*   |   ¡  |  ¡  |  ¡  |  ¡ }|  |¡S )z[Prepare and return the arguments for the ufunc.
        Does not call to_device().
        )r:   r>   rD   rG   rV   )r,   rU   r   r   r   Úget_argumentsÆ   s
    zUFuncMechanism.get_argumentsc                 C   s   | j | j S )z)Returns (result_dtype, function)
        )r'   r)   ©r,   r   r   r   Úget_functionÐ   s    zUFuncMechanism.get_functionc                 C   s   dS )zBIs the `obj` a device array?
        Override in subclass
        Fr   ©r,   Úobjr   r   r   r/   Õ   s    zUFuncMechanism.is_device_arrayc                 C   s   |S )z�Convert the `obj` to a device array
        Override in subclass

        Default implementation is an identity function
        r   rZ   r   r   r   r0   Û   s    zUFuncMechanism.as_device_arrayc                 C   s   t dƒ‚dS )zTHandles ondevice broadcasting

        Override in subclass to add support.
        z'broadcasting on device is not supportedN©ÚNotImplementedError©r,   r=   rI   r   r   r   rP   ã   s    zUFuncMechanism.broadcast_devicec                 C   s   |S )zSEnsures array layout met device requirement.

        Override in sublcass
        r   )r,   r=   r   r   r   rT   ê   s    z!UFuncMechanism.force_array_layoutc                    s  |  d| j¡‰|  dd¡}|r2t dd |¡ ¡ | ||ƒ‰ˆ ¡ }ˆ ¡ \}}|d j}|dk	rvˆ |¡rvˆ 	|¡}‡‡fdd„‰ |d j
d	kr¤‡ fd
d„|D ƒ}g }d}	|D ]6}
ˆ |
¡rÎ| |
¡ d}	q°ˆj|
ˆd�}| |¡ q°|d j}|dk�rLˆj||ˆd�}| |g¡ ˆ ||d ˆ|¡ |	�r<| |¡S | ¡  |¡S n²ˆ |¡�rš|j
d	k�rlˆ |ƒ}|}| |g¡ ˆ ||d ˆ|¡ | |¡S |j|k�sªt‚|j|k�sºt‚ˆj||ˆd�}| |g¡ ˆ ||d ˆ|¡ |j|ˆd� |¡S dS )z1Perform the entire ufunc call mechanism.
        ÚstreamÚoutNzunrecognized keywords: %sú, r   c                    s\   ˆ j r
t‚z
|  ¡ W S  tk
rV   ˆ  | ¡s2‚ n ˆ  | ˆ¡ ¡ }ˆ  |ˆ¡ Y S Y nX d S r   )ÚSUPPORT_DEVICE_SLICINGr]   Zravelr/   Úto_hostÚ	to_device)r   Úhostary)Úcrr_   r   r   Úattempt_ravel  s    

z*UFuncMechanism.call.<locals>.attempt_ravelr   c                    s   g | ]}ˆ |ƒ‘qS r   r   rJ   )rg   r   r   rK     s     z'UFuncMechanism.call.<locals>.<listcomp>FT)r_   )ÚpopÚDEFAULT_STREAMÚwarningsÚwarnÚjoinrW   rY   rI   r/   r0   rL   r7   rd   Úallocate_device_arrayÚextendÚlaunchÚreshapeZcopy_to_hostr    r;   )Úclsr'   r(   Úkwsr`   ZrestyÚfuncZoutshapeZdevarysZ
any_devicer   Zdev_arI   Zdevoutr   )rg   rf   r_   r   Úcallñ   sT    








zUFuncMechanism.callc                 C   s   t ‚dS )zBImplement to device transfer
        Override in subclass
        Nr\   )r,   re   r_   r   r   r   rd   K  s    zUFuncMechanism.to_devicec                 C   s   t ‚dS )z@Implement to host transfer
        Override in subclass
        Nr\   )r,   Údevaryr_   r   r   r   rc   Q  s    zUFuncMechanism.to_hostc                 C   s   t ‚dS )zBImplements device allocation
        Override in subclass
        Nr\   )r,   rI   r;   r_   r   r   r   rm   W  s    z$UFuncMechanism.allocate_device_arrayc                 C   s   t ‚dS )zKImplements device function invocation
        Override in subclass
        Nr\   )r,   rs   Úcountr_   r(   r   r   r   ro   ]  s    zUFuncMechanism.launch)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ri   rb   r.   r:   r>   rD   rG   rV   rW   rY   r/   r0   rP   rT   Úclassmethodrt   rd   rc   rm   ro   r   r   r   r   r&   G   s*   +	!

Yr&   c                 C   s    t | tjƒr| j} t t| ƒ¡S r   )r1   r   Z
EnumMemberr;   r5   Ústr)Útyr   r   r   Úto_dtyped  s    r~   c                   @   sZ   e Zd Zddi fdd„Zedd„ ƒZddd„Zd	d
„ Zdd„ Zdd„ Z	dd„ Z
dd„ ZdS )ÚDeviceVectorizeNFc                 C   s`   |rt dƒ‚|D ]2}|dkr*t dt¡ qd}|d7 }t|| ƒ‚q|| _t|ƒ| _tƒ | _	d S )Núcaching is not supportedÚnopythonz+nopython kwarg for cuda target is redundantzUnrecognized options. z3cuda vectorize target does not support option: '%s')
rB   rj   rk   ÚRuntimeWarningÚKeyErrorÚpy_funcr   Úidentityr   Ú	kernelmap)r,   rs   r…   ÚcacheÚtargetoptionsÚoptÚfmtr   r   r   r.   k  s    ÿ
zDeviceVectorize.__init__c                 C   s   | j S r   ©r„   rX   r   r   r   Úpyfunc{  s    zDeviceVectorize.pyfuncc                 C   sÈ   t  |¡\}}t|f|žŽ }| jj}|  | j||¡}|  |¡\}}|  |¡}tt	j
fdd„ |D ƒ|d d … g žŽ }t||ƒ |d|  }	|  |	|¡}
tdd„ |jD ƒƒ}t|ƒ}||
f| jt|ƒ< d S )Nc                 S   s   g | ]}|d d … ‘qS r   r   rJ   r   r   r   rK   ‰  s     z'DeviceVectorize.add.<locals>.<listcomp>z__vectorized_%sc                 s   s   | ]}t |ƒV  qd S r   )r~   ©r   Útr   r   r   r   �  s     z&DeviceVectorize.add.<locals>.<genexpr>)r	   Únormalize_signaturer
   rŒ   rw   Ú_get_kernel_sourceÚ_kernel_templateÚ_compile_coreÚ_get_globalsr   ÚvoidÚexecÚ_compile_kernelr   r(   r~   r†   )r,   Úsigr(   Úreturn_typeZdevfnsigÚfuncnameZkernelsourceÚcorefnZglblZstagerÚkernelZ	argdtypesZresdtyper   r   r   Úadd  s      ÿ
(
zDeviceVectorize.addc                 C   s   t ‚d S r   r\   rX   r   r   r   Úbuild_ufunc“  s    zDeviceVectorize.build_ufuncc                 C   sH   dd„ t t|jƒƒD ƒ}t|d |¡d dd„ |D ƒ¡d�}|jf |ŽS )Nc                 S   s   g | ]}d | ‘qS )za%dr   ©r   r$   r   r   r   rK   —  s     z6DeviceVectorize._get_kernel_source.<locals>.<listcomp>ra   c                 s   s   | ]}d | V  qdS )z%s[__tid__]Nr   rž   r   r   r   r   š  s     z5DeviceVectorize._get_kernel_source.<locals>.<genexpr>)Únamer(   Úargitems)rQ   r   r(   Údictrl   r   )r,   Útemplater—   r™   r(   Zfmtsr   r   r   r�   –  s    þz"DeviceVectorize._get_kernel_sourcec                 C   s   t ‚d S r   r\   ©r,   r—   r   r   r   r’   �  s    zDeviceVectorize._compile_corec                 C   s   t ‚d S r   r\   )r,   rš   r   r   r   r“      s    zDeviceVectorize._get_globalsc                 C   s   t ‚d S r   r\   ©r,   Úfnobjr—   r   r   r   r–   £  s    zDeviceVectorize._compile_kernel)N)rw   rx   ry   r.   ÚpropertyrŒ   rœ   r�   r�   r’   r“   r–   r   r   r   r   r   j  s   

r   c                   @   sD   e Zd Zddi dfdd„Zedd„ ƒZddd	„Zd
d„ Zdd„ ZdS )ÚDeviceGUFuncVectorizeNFr   c           	      C   sŽ   |rt dƒ‚|rt dƒ‚| dd¡s,t dƒ‚|rZd dd„ | ¡ D ƒ¡}d	}t | |¡ƒ‚|| _t|ƒ| _|| _t	| jƒ\| _
| _tƒ | _d S )
Nr€   zwritable_args are not supportedr�   Tznopython flag must be Truera   c                 S   s   g | ]}t |ƒ‘qS r   )Úrepr©r   Úkr   r   r   rK   ´  s     z2DeviceGUFuncVectorize.__init__.<locals>.<listcomp>z3The following target options are not supported: {0})rB   rh   rl   Úkeysr   r„   r   r…   r
   r   ÚinputsigÚ	outputsigr   r†   )	r,   rs   r—   r…   r‡   rˆ   Zwritable_argsÚoptsrŠ   r   r   r   r.   ¨  s    
zDeviceGUFuncVectorize.__init__c                 C   s   | j S r   r‹   rX   r   r   r   rŒ   À  s    zDeviceGUFuncVectorize.pyfuncc                 C   s  dd„ | j D ƒ}dd„ | jD ƒ}t |¡\}}|tjd fk}|sVtd|› d|› d�ƒ‚| jj}t	| j
||||ƒ}|  |¡}	t||	ƒ |	dj|d� }
tt||| ƒƒ}| j|
t|ƒd	�}t|ƒ}d
d„ |D ƒ}t|d | … ƒ}t|| d … ƒ}||f| j|< d S )Nc                 S   s   g | ]}t |ƒ‘qS r   ©r   ©r   Úxr   r   r   rK   Å  s     z-DeviceGUFuncVectorize.add.<locals>.<listcomp>c                 S   s   g | ]}t |ƒ‘qS r   r¯   r°   r   r   r   rK   Æ  s     z7guvectorized functions cannot return values: signature z specifies z return typez__gufunc_{name})rŸ   )r—   c                 S   s   g | ]}t  t|jƒ¡‘qS r   )r5   r;   r|   r�   r   r   r   rK   Þ  s     )r¬   r­   r	   r�   r   ÚnonerB   r„   rw   Úexpand_gufunc_templater‘   r“   r•   r   rR   Ú_determine_gufunc_outer_typesr–   r   r   r†   )r,   r—   ÚindimsÚoutdimsr(   r˜   Zvalid_return_typer™   ÚsrcZglblsr¥   Zoutertysr›   ÚnoutZdtypesÚindtypesÚ	outdtypesr   r   r   rœ   Ä  s,      ÿ

zDeviceGUFuncVectorize.addc                 C   s   t ‚d S r   r\   r¤   r   r   r   r–   ä  s    z%DeviceGUFuncVectorize._compile_kernelc                 C   s   t ‚d S r   r\   r£   r   r   r   r“   ç  s    z"DeviceGUFuncVectorize._get_globals)N)	rw   rx   ry   r.   r¦   rŒ   rœ   r–   r“   r   r   r   r   r§   §  s   ÿ


 r§   c                 c   sZ   t | |ƒD ]J\}}t|tjƒr2|j|d d�V  q
|dkrBtdƒ‚tj|ddd�V  q
d S )Nr   )rL   r   z,gufunc signature mismatch: ndim>0 for scalarÚA)r;   rL   Zlayout)r   r1   r   ÚArrayÚcopyr   )ZargtysZdimsÚatÚndr   r   r   r´   ë  s    r´   c                 C   s¦   || }dd„ t t|ƒƒD ƒ}d d dd„ |D ƒ¡¡}dd„ t|||ƒD ƒ}dd„ t|t|ƒd… ||t|ƒd… ƒD ƒ}	||	 }
| j|d |¡|d |
¡d	�}|S )
z"Expand gufunc source template
    c                 S   s   g | ]}d   |¡‘qS )zarg{0}©r   rž   r   r   r   rK   ù  s     z*expand_gufunc_template.<locals>.<listcomp>zmin({0})ra   c                 S   s   g | ]}d   |¡‘qS )z{0}.shape[0]rÀ   rJ   r   r   r   rK   ú  s   ÿc                 S   s   g | ]\}}}t |||ƒ‘qS r   ©Ú_gen_src_for_indexing©r   ÚarefÚadimsÚatyper   r   r   rK   ü  s   ÿc                 S   s   g | ]\}}}t |||ƒ‘qS r   rÁ   rÃ   r   r   r   rK   þ  s   ÿN)rŸ   r(   Ú
checkedargr    )rQ   r   r   rl   r   )r¢   rµ   r¶   r™   r)   ZargdimsÚargnamesrÇ   ÚinputsÚoutputsr    r·   r   r   r   r³   õ  s&    ÿ

ÿÿÿþr³   c                 C   s   dj | t||ƒd�S )Nz{aref}[{sliced}])rÄ   Zsliced)r   Ú_gen_src_index)rÄ   rÅ   rÆ   r   r   r   rÂ     s    ÿrÂ   c                 C   sD   | dkrd  dgdg|   ¡S t|tjƒr<|jd | kr<dS dS d S )Nr   ú,Z__tid__ú:r   z__tid__:(__tid__ + 1))rl   r1   r   r¼   rL   )rÅ   rÆ   r   r   r   rË     s
    rË   c                   @   s,   e Zd ZdZedd„ ƒZdd„ Zdd„ ZdS )	ÚGUFuncEnginezZDetermine how to broadcast and execute a gufunc
    base on input shape and signature
    c                 C   s   | t |ƒŽ S r   r   )rq   r
   r   r   r   Úfrom_signature  s    zGUFuncEngine.from_signaturec                 C   s(   || _ || _t| j ƒ| _t| jƒ| _d S r   )ÚsinÚsoutr   Úninr¸   )r,   r¬   r­   r   r   r   r.   "  s    zGUFuncEngine.__init__c                 C   sÐ  t |ƒ| jkrtdƒ‚i }g }g }tt|| jƒƒD ]Ð\}\}}|d7 }t |ƒ}t |ƒ|k rld}	t|	|f ƒ‚|rŽ|| d … }
|d | … }nd}
|}tt|
|ƒƒD ]H\}\}}|t |ƒ7 }||krä|| |kräd}	t|	||f ƒ‚|||< q¤| |¡ | |
¡ q2g }| jD ]2}g }|D ]}| || ¡ �q| t	|ƒ¡ �qdd„ |D ƒ}t
 |¡}|| }dg| j }t|ƒD ]H\}}||k�rv|d	k�sœ|dk�r¦d
||< nd}	t|	|d f ƒ‚�qvt| ||||ƒS )Nz invalid number of input argumentr   z%arg #%d: insufficient inner dimensionr   z$arg #%d: shape[%d] mismatch argumentc                 S   s   g | ]}t tj|d ƒ‘qS r   )r   ÚoperatorÚmul©r   Úsr   r   r   rK   T  s     z)GUFuncEngine.schedule.<locals>.<listcomp>Fr   Tz!arg #%d: outer dimension mismatch)r   rÒ   rB   r!   r   rÐ   r   r7   rÑ   r   r5   ZargmaxÚGUFuncSchedule)r,   ÚishapesZ	symbolmapZouter_shapesZinner_shapesZargnrI   ÚsymbolsZ
inner_ndimrŠ   Zinner_shapeZouter_shapeZaxisÚdimÚsymÚoshapesZoutsigÚoshapeÚsizesZ	largest_iÚloopdimsÚpinnedr$   Údr   r   r   Úschedule*  sT    





zGUFuncEngine.scheduleN)rw   rx   ry   rz   r{   rÏ   r.   râ   r   r   r   r   rÎ     s
   
rÎ   c                   @   s   e Zd Zdd„ Zdd„ ZdS )r×   c                    sF   || _ || _|| _ˆ | _ttjˆ dƒ| _|| _‡ fdd„|D ƒ| _	d S )Nr   c                    s   g | ]}ˆ | ‘qS r   r   rÕ   ©rß   r   r   rK   p  s     z+GUFuncSchedule.__init__.<locals>.<listcomp>)
ÚparentrØ   rÜ   rß   r   rÓ   rÔ   Úloopnrà   Úoutput_shapes)r,   rä   rØ   rÜ   rß   rà   r   rã   r   r.   e  s    zGUFuncSchedule.__init__c                    s,   dd l }d}‡ fdd„|D ƒ}| t|ƒ¡S )Nr   )rØ   rÜ   rß   rå   rà   c                    s   g | ]}|t ˆ |ƒf‘qS r   )r<   r©   rX   r   r   rK   v  s     z*GUFuncSchedule.__str__.<locals>.<listcomp>)ÚpprintÚpformatr¡   )r,   rç   ÚattrsÚvaluesr   rX   r   Ú__str__r  s    zGUFuncSchedule.__str__N)rw   rx   ry   r.   rë   r   r   r   r   r×   d  s   r×   c                   @   sL   e 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d„ Z
dS )ÚGeneralizedUFuncc                 C   s   || _ || _d| _d S )Ni   @)r†   ÚengineZmax_blocksize)r,   r†   rí   r   r   r   r.   {  s    zGeneralizedUFunc.__init__c                 O   sv   |   | jj| jj||¡}|  |j|j¡\}}}}| |¡ | ||¡}| 	¡ }	|  
||	|¡}
| ||j|
¡ | |¡S r   )Z_call_stepsrí   rÒ   r¸   Ú	_schedulerÉ   rÊ   Úadjust_input_typesÚprepare_outputsÚprepare_inputsrV   Úlaunch_kernelrå   Úpost_process_outputs)r,   r(   rr   Z	callstepsr¹   râ   rº   r›   rÊ   rÉ   Ú
parametersr   r   r   Ú__call__€  s     ÿ ÿ
zGeneralizedUFunc.__call__c           
      C   s¨   dd„ |D ƒ}| j  |¡}tdd„ |D ƒƒ}z| j| \}}W n, tk
rj   |  |¡}| j| \}}Y nX t|j|ƒD ]"\}}	|	d k	rx||	jkrxt	dƒ‚qx||||fS )Nc                 S   s   g | ]
}|j ‘qS r   rH   rJ   r   r   r   rK   �  s     z.GeneralizedUFunc._schedule.<locals>.<listcomp>c                 s   s   | ]}|j V  qd S r   rE   rž   r   r   r   r   ”  s     z-GeneralizedUFunc._schedule.<locals>.<genexpr>zoutput shape mismatch)
rí   râ   r   r†   rƒ   Ú_search_matching_signaturer   ræ   rI   r   )
r,   rÉ   ZoutsZinput_shapesrâ   r¹   rº   r›   Zsched_shaper`   r   r   r   rî   �  s    

zGeneralizedUFunc._schedulec                 C   s<   | j  ¡ D ]$}tdd„ t||ƒD ƒƒr
|  S q
tdƒ‚dS )zÌ
        Given the input types in `idtypes`, return a compatible sequence of
        types that is defined in `kernelmap`.

        Note: Ordering is guaranteed by `kernelmap` being a OrderedDict
        c                 s   s   | ]\}}t  ||¡V  qd S r   )r5   Zcan_cast)r   r@   Zdesiredr   r   r   r   ®  s   ÿz>GeneralizedUFunc._search_matching_signature.<locals>.<genexpr>zno matching signatureN)r†   r«   rA   r   rB   )r,   Zidtypesr—   r   r   r   rö   ¦  s    ÿ
z+GeneralizedUFunc._search_matching_signaturec                 C   s¶   |j dkstdƒ‚|jsdn|j }g }t||jƒD ]B\}}|s`|jdkr`|  ||¡}| |¡ q2| |  |||¡¡ q2g }	t||j	ƒD ]\}
}|	 |
j
|f|žŽ ¡ q†t|ƒt|	ƒ S )Nr   zzero looping dimensionr   )rå   r    rß   r   rØ   ÚsizeÚ_broadcast_scalar_inputr7   Ú_broadcast_arrayrÜ   rp   r   )r,   râ   ÚparamsZretvalsZodimÚ	newparamsÚpÚcsru   Z
newretvalsÚretvalrÝ   r   r   r   rV   ´  s    zGeneralizedUFunc._broadcastc                 C   sf   |f| }|j |kr|S t|j ƒt|ƒk rX|t|j ƒ d … |j ksLtdƒ‚|  ||¡S |j|Ž S d S )Nz+cannot add dim and reshape at the same time)rI   r   r    Ú_broadcast_add_axisrp   )r,   r=   ZnewdimZinnerdimÚnewshaper   r   r   rù   Ç  s    

ÿz!GeneralizedUFunc._broadcast_arrayc                 C   s   t dƒ‚d S )Nzcannot add new axisr\   )r,   r=   r   r   r   r   rÿ   ×  s    z$GeneralizedUFunc._broadcast_add_axisc                 C   s   t ‚d S r   r\   r^   r   r   r   rø   Ú  s    z(GeneralizedUFunc._broadcast_scalar_inputN)rw   rx   ry   r.   rõ   rî   rö   rV   rù   rÿ   rø   r   r   r   r   rì   z  s   rì   c                   @   s~   e Zd ZdZdddgZedd„ ƒZedd„ ƒZed	d
„ ƒZedd„ ƒZ	edd„ ƒZ
dd„ Zdd„ Zdd„ Zdd„ Zdd„ ZdS )ÚGUFuncCallStepsab  
    Implements memory management and kernel launch operations for GUFunc calls.

    One instance of this class is instantiated for each call, and the instance
    is specific to the arguments given to the GUFunc call.

    The base class implements the overall logic; subclasses provide
    target-specific implementations of individual functions.
    rÊ   rÉ   Ú_copy_result_to_hostc                 C   s   dS )zImplement the kernel launchNr   )r,   r›   Znelemr(   r   r   r   rò   ð  s    zGUFuncCallSteps.launch_kernelc                 C   s   dS )zb
        Return True if `obj` is a device array for this target, False
        otherwise.
        Nr   rZ   r   r   r   r/   ô  s    zGUFuncCallSteps.is_device_arrayc                 C   s   dS )zƒ
        Return `obj` as a device array on this target.

        May return `obj` directly if it is already on the target.
        Nr   rZ   r   r   r   r0   û  s    zGUFuncCallSteps.as_device_arrayc                 C   s   dS )zK
        Copy `hostary` to the device and return the device array.
        Nr   )r,   re   r   r   r   rd     s    zGUFuncCallSteps.to_devicec                 C   s   dS )zc
        Allocate a new uninitialized device array with the given shape and
        dtype.
        Nr   )r,   rI   r;   r   r   r   rm   	  s    z%GUFuncCallSteps.allocate_device_arrayc                    s6  |  d¡}|d krbt|ƒ||| fkrbdd„ }d||ƒ› d||| ƒ› d|t|ƒƒ› d�}t|ƒ‚|d k	r€t|ƒ|kr€tdƒ‚n
|g| }d	}g ˆ_|D ]2}	ˆ |	¡r¾ˆj ˆ |	¡¡ d
}q˜ˆj |	¡ q˜t‡fdd„|D ƒƒ }
|
oê|ˆ_	‡fdd„‰ ‡ fdd„|D ƒ}|d |… ˆ_
||d … }|�r2|ˆ_d S )Nr`   c                 S   s   | › dd| dk › �S )Nz positional argumentrÖ   r   r   )Únr   r   r   Úpos_argn  s    z*GUFuncCallSteps.__init__.<locals>.pos_argnzThis gufunc accepts z  (when providing input only) or z( (when providing input and output). Got Ú.z<cannot specify argument 'out' as both positional and keywordTFc                    s   g | ]}ˆ   |¡‘qS r   )r/   rJ   rX   r   r   rK   2  s     z,GUFuncCallSteps.__init__.<locals>.<listcomp>c                    s    ˆ   | ¡rˆ j}ntj}|| ƒS r   )r/   r0   r5   r8   )r   ÚconvertrX   r   r   Únormalize_arg;  s    
z/GUFuncCallSteps.__init__.<locals>.normalize_argc                    s   g | ]}ˆ |ƒ‘qS r   r   rJ   )r  r   r   rK   C  s     )Úgetr   rB   r   rÊ   r/   r7   r0   Úanyr  rÉ   )r,   rÒ   r¸   r(   ÚkwargsrÊ   r  ÚmsgZall_user_outputs_are_hostÚoutputZall_host_arraysZnormalized_argsZunused_inputsr   )r  r,   r   r.     s2    
,


ÿzGUFuncCallSteps.__init__c                 C   s\   t t|| jƒƒD ]F\}\}}||jkrt|dƒsFd t|ƒ¡}t|ƒ‚| |¡| j|< qdS )zØ
        Attempt to cast the inputs to the required types if necessary
        and if they are not device arrays.

        Side effect: Only affects the elements of `inputs` that require
        a type cast.
        ÚastypezNcompatible signature is possible by casting but {0} does not support .astype()N)	r!   r   rÉ   r;   Úhasattrr   ÚtyperB   r  )r,   r¹   r$   ZityÚvalr  r   r   r   rï   K  s    

ÿz"GUFuncCallSteps.adjust_input_typesc                 C   sH   g }t |j|| jƒD ].\}}}|dks,| jr8|  ||¡}| |¡ q|S )zç
        Returns a list of output parameters that all reside on the target device.

        Outputs that were passed-in to the GUFunc are used if they reside on the
        device; other outputs are allocated as necessary.
        N)r   ræ   rÊ   r  rm   r7   )r,   râ   rº   rÊ   rI   r;   r  r   r   r   rð   \  s    ÿzGUFuncCallSteps.prepare_outputsc                    s    ‡fdd„‰ ‡ fdd„ˆj D ƒS )zZ
        Returns a list of input parameters that all reside on the target device.
        c                    s    ˆ   | ¡rˆ j}nˆ j}|| ƒS r   )r/   r0   rd   )Z	parameterr  rX   r   r   Úensure_devicep  s    
z5GUFuncCallSteps.prepare_inputs.<locals>.ensure_devicec                    s   g | ]}ˆ |ƒ‘qS r   r   )r   rü   )r  r   r   rK   x  s     z2GUFuncCallSteps.prepare_inputs.<locals>.<listcomp>)rÉ   rX   r   )r  r,   r   rñ   l  s    zGUFuncCallSteps.prepare_inputsc                    sV   ˆ j r"‡ fdd„t|ˆ jƒD ƒ}nˆ jd dk	r6ˆ j}t|ƒdkrJ|d S t|ƒS dS )a+  
        Moves the given output(s) to the host if necessary.

        Returns a single value (e.g. an array) if there was one output, or a
        tuple of arrays if there were multiple. Although this feels a little
        jarring, it is consistent with the behavior of GUFuncs in general.
        c                    s   g | ]\}}ˆ   ||¡‘qS r   )rc   )r   r  Zself_outputrX   r   r   rK   ƒ  s   ÿz8GUFuncCallSteps.post_process_outputs.<locals>.<listcomp>r   Nr   )r  r   rÊ   r   r   )r,   rÊ   r   rX   r   ró   z  s    

ÿz$GUFuncCallSteps.post_process_outputsN)rw   rx   ry   rz   Ú	__slots__r   rò   r/   r0   rd   rm   r.   rï   rð   rñ   ró   r   r   r   r   r  Þ  s(   ý




;r  )Ú	metaclass)&rz   Úabcr   r   Úcollectionsr   rÓ   rj   Ú	functoolsr   Únumpyr5   Znumba.np.ufunc.ufuncbuilderr   r   Z
numba.corer   r	   Znumba.core.typingr
   Znumba.np.ufunc.sigparser   r   r   r%   Úobjectr&   r~   r   r§   r´   r³   rÂ   rË   rÎ   r×   rì   r  r   r   r   r   Ú<module>   s6     =D
Kd