U
    hâËd\0  ã                   @   sŠ  d dl mZ d dlmZmZmZmZmZmZ d dl	m
Z
mZmZmZmZmZ d dlmZ d dlmZmZmZm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! d
d„ Z"G dd„ deƒZ#G dd„ deƒZ$dd„ Z%eddd�G dd„ deƒƒZ&eddd�G dd„ deƒƒZ'eddd�G dd„ deƒƒZ(G dd„ deƒZ)ed*dd„ƒZ*ed+d d!„ƒZ+d,d"d#„Z,d$d%„ Z-d&d'„ Z.G d(d)„ d)e/ƒZ0dS )-é    )ÚConcreteTemplate)ÚtypesÚtypingÚfuncdescÚconfigÚcompilerÚsigutils)Úsanitize_compile_result_entriesÚCompilerBaseÚDefaultPassBuilderÚFlagsÚOptionÚCompileResult)Úglobal_compiler_lock)ÚLoweringPassÚAnalysisPassÚPassManagerÚregister_pass)ÚNumbaInvalidConfigWarningÚTypingError)ÚIRLegalizationÚNativeLoweringÚAnnotateTypes)Úwarn)Úget_current_devicec                 C   s"   | d krd S t | tƒst‚| S d S ©N)Ú
isinstanceÚdictÚAssertionError)Úx© r    úL/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numba/cuda/compiler.pyÚ_nvvm_options_type   s    r"   c                   @   s(   e Zd Zeeddd�Zeeddd�ZdS )Ú	CUDAFlagsNzNVVM options)ÚtypeÚdefaultÚdoczCompute Capability)Ú__name__Ú
__module__Ú__qualname__r   r"   Únvvm_optionsÚtupleÚcompute_capabilityr    r    r    r!   r#      s   ýýr#   c                   @   s   e Zd Zedd„ ƒZdS )ÚCUDACompileResultc                 C   s   t | ƒS r   )Úid©Úselfr    r    r!   Úentry_point7   s    zCUDACompileResult.entry_pointN)r'   r(   r)   Úpropertyr1   r    r    r    r!   r-   6   s   r-   c                  K   s   t | ƒ} tf | ŽS r   )r	   r-   )Úentriesr    r    r!   Úcuda_compile_result<   s    r4   TF)Zmutates_CFGZanalysis_onlyc                   @   s    e Zd ZdZdd„ Zdd„ ZdS )ÚCUDABackendZcuda_backendc                 C   s   t  | ¡ d S r   ©r   Ú__init__r/   r    r    r!   r7   F   s    zCUDABackend.__init__c              
   C   sJ   |d }t j|jf|jžŽ }t|j|j|jj|j	|j
|j||jd�|_dS )zH
        Back-end: Packages lowering output in a compile result
        Úcr)Útyping_contextÚtarget_contextZtyping_errorÚtype_annotationÚlibraryÚcall_helperÚ	signatureÚfndescT)r   r>   Úreturn_typeÚargsr4   Ú	typingctxÚ	targetctxÚstatusZfail_reasonr;   r<   r=   r?   r8   )r0   ÚstateZloweredr>   r    r    r!   Úrun_passI   s    ø
zCUDABackend.run_passN©r'   r(   r)   Ú_namer7   rF   r    r    r    r!   r5   A   s   r5   c                   @   s$   e Zd ZdZdZdd„ Zdd„ ZdS )ÚCreateLibraryzå
    Create a CUDACodeLibrary for the NativeLowering pass to populate. The
    NativeLowering pass will create a code library if none exists, but we need
    to set it up with nvvm_options from the flags if they are present.
    Úcreate_libraryc                 C   s   t  | ¡ d S r   r6   r/   r    r    r!   r7   g   s    zCreateLibrary.__init__c                 C   s8   |j  ¡ }|jj}|jj}|j||d�|_|j ¡  dS )N)r*   T)	rC   ÚcodegenZfunc_idZfunc_qualnameÚflagsr*   rJ   r<   Zenable_object_caching)r0   rE   rK   Únamer*   r    r    r!   rF   j   s    

zCreateLibrary.run_passN)r'   r(   r)   Ú__doc__rH   r7   rF   r    r    r    r!   rI   ]   s   rI   c                   @   s    e Zd ZdZdd„ Zdd„ ZdS )ÚCUDALegalizationZcuda_legalizationc                 C   s   t  | ¡ d S r   )r   r7   r/   r    r    r!   r7   z   s    zCUDALegalization.__init__c                    sX   ddl m} |ƒ jrdS |j}‡ ‡fdd„‰ | ¡ D ]\‰}t|tjƒr4ˆ |jƒ q4dS )Nr   )ÚNVVMFc                    sT   t | tjtjfƒr&ˆ› d�}t|ƒ‚n*t | tjƒrP| j ¡ D ]}ˆ |d jƒ q<d S )Nzæ is a char sequence type. This type is not supported with CUDA toolkit versions < 11.2. To use this type, you need to update your CUDA toolkit - try 'conda install cudatoolkit=11' if you are using conda to manage your environment.é   )	r   r   ZUnicodeCharSeqZCharSeqr   ZRecordÚfieldsÚitemsr$   )ÚdtypeÚmsgZsubdtype©Úcheck_dtypeÚkr    r!   rW   …   s    

z.CUDALegalization.run_pass.<locals>.check_dtype)	Znumba.cuda.cudadrv.nvvmrP   Z	is_nvvm70ÚtypemaprS   r   r   ZArrayrT   )r0   rE   rP   ZtypmapÚvr    rV   r!   rF   }   s    zCUDALegalization.run_passNrG   r    r    r    r!   rO   u   s   rO   c                   @   s   e Zd Zdd„ Zdd„ ZdS )ÚCUDACompilerc                 C   st   t }tdƒ}| | j¡}|j |j¡ | | j¡}|j |j¡ | td¡ |  	| j¡}|j |j¡ | 
¡  |gS )NÚcudazCUDA legalization)r   r   Zdefine_untyped_pipelinerE   ZpassesÚextendZdefine_typed_pipelineÚadd_passrO   Údefine_cuda_lowering_pipelineÚfinalize)r0   ZdpbÚpmZuntyped_passesZtyped_passesZlowering_passesr    r    r!   Údefine_pipelines™   s    zCUDACompiler.define_pipelinesc                 C   sP   t dƒ}| td¡ | td¡ | td¡ | td¡ | td¡ | ¡  |S )NZcuda_loweringz$ensure IR is legal prior to loweringzannotate typeszcreate libraryznative loweringzcuda backend)r   r^   r   r   rI   r   r5   r`   )r0   rE   ra   r    r    r!   r_   ª   s    ÿz*CUDACompiler.define_cuda_lowering_pipelineN)r'   r(   r)   rb   r_   r    r    r    r!   r[   ˜   s   r[   Nc	                 C   sÚ   |d krt dƒ‚ddlm}	 |	j}
|	j}tƒ }d|_d|_d|_|sH|rNd|_	|rXd|_
|rdd|_nd|_|rtd|_|r~d|_|rˆ||_||_ddlm} |d	ƒ�  tj|
|| |||i td
�}W 5 Q R X |j}| ¡  |S )Nz#Compute Capability must be suppliedrQ   ©Úcuda_targetTÚpythonÚnumpyr   )Útarget_overrider\   )rB   rC   ÚfuncrA   r@   rL   ÚlocalsZpipeline_class)Ú
ValueErrorÚ
descriptorrd   r9   r:   r#   Z
no_compileZno_cpython_wrapperZno_cfunc_wrapperZ	debuginfoZdbg_directives_onlyZerror_modelZforceinlineÚfastmathr*   r,   Znumba.core.target_extensionrg   r   Zcompile_extrar[   r<   r`   )Úpyfuncr@   rA   ÚdebugÚlineinfoÚinlinerl   r*   Úccrd   rB   rC   rL   rg   Úcresr<   r    r    r!   Úcompile_cudaº   sJ    
ù	rs   c              
   C   sÒ   |r|rd}t t|ƒƒ ||r"dnddœ}	t |¡\}
}|p@tj}t| ||
||||	|d�}|jj}|r||s||t	j
kr|tdƒ‚|rˆ|j}n6|j}| j}|j}|j}| |j|j|||	||¡\}}|j|d�}||fS )a  Compile a Python function to PTX for a given set of argument types.

    :param pyfunc: The Python function to compile.
    :param sig: The signature representing the function's input and output
                types.
    :param debug: Whether to include debug info in the generated PTX.
    :type debug: bool
    :param lineinfo: Whether to include a line mapping from the generated PTX
                     to the source code. Usually this is used with optimized
                     code (since debug mode would automatically include this),
                     so we want debug info in the LLVM but only the line
                     mapping in the final PTX.
    :type lineinfo: bool
    :param device: Whether to compile a device function. Defaults to ``False``,
                   to compile global kernel functions.
    :type device: bool
    :param fastmath: Whether to enable fast math flags (ftz=1, prec_sqrt=0,
                     prec_div=, and fma=1)
    :type fastmath: bool
    :param cc: Compute capability to compile for, as a tuple
               ``(MAJOR, MINOR)``. Defaults to ``(5, 0)``.
    :type cc: tuple
    :param opt: Enable optimizations. Defaults to ``True``.
    :type opt: bool
    :return: (ptx, resty): The PTX code and inferred return type
    :rtype: tuple
    z{debug=True with opt=True (the default) is not supported by CUDA. This may result in a crash - set debug=False or opt=False.é   r   )rl   Úopt)rn   ro   rl   r*   rq   z'CUDA kernel must have void return type.)rq   )r   r   r   Znormalize_signaturer   ZCUDA_DEFAULT_PTX_CCrs   r>   r@   r   ÚvoidÚ	TypeErrorr<   r:   Ú__code__Úco_filenameÚco_firstlinenoZprepare_cuda_kernelr?   Zget_asm_str)rm   Úsigrn   ro   Údevicerl   rq   ru   rU   r*   rA   r@   rr   ZrestyÚlibZtgtÚcodeÚfilenameZlinenumZkernelZptxr    r    r!   Úcompile_ptxõ   s>    
þ

  þ  þr€   c              
   C   s    t ƒ j}t| ||||||dd�S )zÑCompile a Python function to PTX for a given set of argument types for
    the current device's compute capabilility. This calls :func:`compile_ptx`
    with an appropriate ``cc`` value for the current device.T)rn   ro   r|   rl   rq   ru   )r   r,   r€   )rm   r{   rn   ro   r|   rl   ru   rq   r    r    r!   Úcompile_ptx_for_current_device9  s    
   ÿr�   c                 C   s   t | ||ƒjS r   )Ú declare_device_function_templateÚkey©rM   ÚrestypeÚargtypesr    r    r!   Údeclare_device_functionC  s    r‡   c                    st   ddl m} |j}|j}tj|f|žŽ ‰t| ˆƒ‰ G ‡ ‡fdd„dtƒ}tj	| ||d�}| 
ˆ |¡ | 
ˆ |¡ |S )NrQ   rc   c                       s   e Zd Z” Z”gZdS )zBdeclare_device_function_template.<locals>.device_function_templateN)r'   r(   r)   rƒ   Zcasesr    ©Zextfnr{   r    r!   Údevice_function_templateN  s   r‰   r„   )rk   rd   r9   r:   r   r>   ÚExternFunctionr   r   ZExternalFunctionDescriptorZinsert_user_function)rM   r…   r†   rd   rB   rC   r‰   r?   r    rˆ   r!   r‚   G  s    
  ÿr‚   c                   @   s   e Zd Zdd„ ZdS )rŠ   c                 C   s   || _ || _d S r   )rM   r{   )r0   rM   r{   r    r    r!   r7   [  s    zExternFunction.__init__N)r'   r(   r)   r7   r    r    r    r!   rŠ   Z  s   rŠ   )FFFFNN)FFFFNT)FFFFT)1Znumba.core.typing.templatesr   Z
numba.corer   r   r   r   r   r   Znumba.core.compilerr	   r
   r   r   r   r   Znumba.core.compiler_lockr   Znumba.core.compiler_machineryr   r   r   r   Znumba.core.errorsr   r   Znumba.core.typed_passesr   r   r   Úwarningsr   Znumba.cuda.apir   r"   r#   r-   r4   r5   rI   rO   r[   rs   r€   r�   r‡   r‚   ÚobjectrŠ   r    r    r    r!   Ú<module>   sP     	


""       þ:      ÿC      ÿ

