Ë
    óÍ:j¯0  ã                   óR  — d dl Z d dlZd dlmZmZmZ d dlZd dlmZ	 de j                  fd„Zdeddfd„Zde j                  fd„Z	 	 	 	 dd	ed
edee   dedee   dee   defd„Z G d„ d«      Z G d„ d«      Z	 ddeeef   deee      deeeedf   f   fd„Z	 ddedededefd„Zy)é    N)ÚAnyÚOptionalÚUnion)Ú_get_device_indexÚreturnc                  ó|   — t         j                  dk(  rt        j                  d«      S t        j                  d«      S )NÚwin32z
nvcuda.dllzlibcuda.so.1©ÚsysÚplatformÚctypesÚCDLL© ó    úf/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torch/cuda/_utils.pyÚ_get_cuda_libraryr      s,   € Ü
‡|�|�wÒÜ�{‰{˜<Ó(Ð(ä�{‰{˜>Ó*Ð*r   Úresultc                 ó   — | dk(  ry t        j                  «       }t        «       }|j                  | t        j                  |«      «       |j
                  �|j
                  j                  «       nd}t        d|› �«      ‚)Nr   úUnknown CUDA errorúCUDA error: )r   Úc_char_pr   ÚcuGetErrorStringÚbyrefÚvalueÚdecodeÚRuntimeError)r   Úerr_strÚlibcudaÚerror_messages       r   Ú_check_cudar       sn   € Ø�‚{ØÜ�o‰oÓ€GÜÓ!€GØ×Ñ˜V¤V§\¡\°'Ó%:Ô;à")§-¡-Ð";ˆ�‰×ÑÔÐAUð ô ˜ m _Ð5Ó
6Ð6r   c                  ó|   — t         j                  dk(  rt        j                  d«      S t        j                  d«      S )Nr	   znvrtc64_120_0.dllzlibnvrtc.sor
   r   r   r   Ú_get_nvrtc_libraryr"       s/   € ô ‡|�|�wÒÜ�{‰{Ð.Ó/Ð/ä�{‰{˜=Ó)Ð)r   Úkernel_sourceÚkernel_nameÚcompute_capabilityÚheader_codeÚcuda_include_dirsÚnvcc_optionsc           
      ó€  ‡‡— ddl }t        «       ŠdŠdt        ddfˆˆfd„}| j                  «       j	                  d«      sd| › �} |r	|dz   | z   }n| }|j                  d	«      }	|€M|j                  j                  |j                  j                  «       «      }
|
j                  › |
j                  › �}g }|j                  d
|› �j                  «       «       |r)|D ]$  }|j                  d|› �j                  «       «       Œ& |r'|D ]"  }|j                  |j                  d	«      «       Œ$ ddlm} |D �cg c]
  }|dk7  sŒ	|‘Œ }}|j                  |D �cg c]  }|j                  d	«      ‘Œ c}«       t        |«      }t!        j"                  |z  |Ž }t!        j$                  «       } |‰j'                  t!        j(                  |«      |	|› d�j                  «       ddd«      «       ‰j+                  |||«      }|‰k7  r�t!        j,                  «       }‰j/                  |t!        j(                  |«      «       t!        j0                  |j2                  «      }‰j5                  ||«       t7        d|j2                  j9                  «       › �«      ‚t!        j,                  «       } |‰j;                  |t!        j(                  |«      «      «       t!        j0                  |j2                  «      } |‰j=                  ||«      «       ‰j?                  t!        j(                  |«      «       |j2                  S c c}w c c}w )a°  
    Compiles a CUDA kernel using NVRTC and returns the PTX code.

    Args:
        kernel_source (str): The CUDA kernel source code as a string
        kernel_name (str): The name of the kernel function to compile
        compute_capability (str, None): The compute capability to target (e.g., "86").
                                           If None, will detect from current device.
        header_code (str, optional): Additional header code to prepend to the kernel source
        cuda_include_dirs (list, None): List of directories containing CUDA headers
        nvcc_options (list, None): Additional options to pass to NVRTC

    Returns:
        str: The compiled PTX code
    r   Nr   r   c                 óî   •— | ‰k7  rot        j                  «       }‰j                  | t        j                  |«      «       |j                  �|j                  j                  «       nd}t        d|› �«      ‚y )Nr   r   )r   r   ÚnvrtcGetErrorStringr   r   r   r   )r   r   r   ÚNVRTC_SUCCESSÚlibnvrtcs      €€r   Úcheck_nvrtcz#_nvrtc_compile.<locals>.check_nvrtcJ   so   ø€ Ø�]Ò"Ü—o‘oÓ'ˆGØ×(Ñ(¨´·±¸gÓ1FÔGð —=‘=Ð,ð —‘×$Ñ$Ô&à)ð ô
  ¨m¨_Ð=Ó>Ð>ð #r   z
extern "C"zextern "C" ú
úutf-8z--gpu-architecture=sm_z-I)ÚCOMMON_NVCC_FLAGSz--expt-relaxed-constexprz.cuzKernel compilation failed:
) Ú
torch.cudar"   ÚintÚstripÚ
startswithÚencodeÚcudaÚget_device_propertiesÚcurrent_deviceÚmajorÚminorÚappendÚtorch.utils.cpp_extensionr1   ÚextendÚlenr   r   Úc_void_pÚnvrtcCreateProgramr   ÚnvrtcCompileProgramÚc_size_tÚnvrtcGetProgramLogSizeÚcreate_string_bufferr   ÚnvrtcGetProgramLogr   r   ÚnvrtcGetPTXSizeÚnvrtcGetPTXÚnvrtcDestroyProgram)r#   r$   r%   r&   r'   r(   Útorchr.   Úfull_sourceÚsource_bytesÚpropsÚoptionsÚ	directoryÚoptionr1   ÚflagÚnvrtc_compatible_flagsÚnum_optionsÚoptions_arrayÚprogÚresÚlog_sizeÚlogÚptx_sizeÚptxr,   r-   s                            @@r   Ú_nvrtc_compiler[   )   s  ù€ ó0 ô "Ó#€Hð €Mð	?œCð 	? Dö 	?ð ×ÑÓ ×+Ñ+¨LÔ9Ø% m _Ð5ˆñ Ø! DÑ(¨=Ñ8‰à#ˆð ×%Ñ% gÓ.€Lð Ð!Ø—
‘
×0Ñ0°·±×1JÑ1JÓ1LÓMˆØ %§¡˜}¨U¯[©[¨MÐ:Ðð €GØ‡N�NÐ+Ð,>Ð+?Ð@×GÑGÓIÔJñ Ø*ò 	6ˆIØ�N‰N˜R 	˜{Ð+×2Ñ2Ó4Õ5ð	6ñ Ø"ò 	3ˆFØ�N‰N˜6Ÿ=™=¨Ó1Õ2ð	3õ <ð +öØ¨dÐ6PÓ.PŠðÐð ð ‡N�NÐ5KÖL¨T�D—K‘K Õ(ÒLÔMô �g“,€KÜ—_‘_ {Ñ2°WÐ=€Mô �?‰?Ó€DÙØ×#Ñ#Ü�L‰L˜ÓØØˆm˜3Ð×&Ñ&Ó(ØØØó	
ô	ð ×
&Ñ
& t¨[¸-Ó
H€Cð ˆmÒä—?‘?Ó$ˆØ×'Ñ'¨¬f¯l©l¸8Ó.DÔEÜ×)Ñ)¨(¯.©.Ó9ˆØ×#Ñ# D¨#Ô.ÜÐ9¸#¿)¹)×:JÑ:JÓ:LÐ9MÐNÓOÐOô �‰Ó €HÙ�×(Ñ(¨¬v¯|©|¸HÓ/EÓFÔGÜ
×
%Ñ
% h§n¡nÓ
5€CÙ�×$Ñ$ T¨3Ó/Ô0Ø× Ñ ¤§¡¨dÓ!3Ô4à�9‰9ÐùòSùò Ms   Ä5
L6Å L6ÅL;c                   ó@   — e Zd Zdej                  ddfd„Zdeddfd„Zy)Ú_CudaModuleÚmoduler   Nc                 ó    — || _         i | _        y ©N)Ú_moduleÚ_kernels)Úselfr^   s     r   Ú__init__z_CudaModule.__init__¦   s   € ØˆŒØ02ˆ�r   ÚnameÚ_CudaKernelc           	      ó   — || j                   v r| j                   |   S ddlm}  |«       }t        j                  «       }	 t        |j                  t        j                  |«      | j                  |j                  d«      «      «       t        || j                  «      }|| j                   |<   |S # t        $ r}t        d|› d�«      |‚d }~ww xY w)Nr   )r   r0   zNo kernel named 'z' in this module)rb   Útorch.cuda._utilsr   r   r@   r    ÚcuModuleGetFunctionr   ra   r6   rf   r   ÚAttributeError)rc   re   r   r   ÚfuncÚkernelÚerrs          r   Ú__getattr__z_CudaModule.__getattr__ª   s¿   € Ø�4—=‘=Ñ Ø—=‘= Ñ&Ð&õ 	8á#Ó%ˆä�‰Ó ˆð	VÜØ×+Ñ+Ü—L‘L Ó&¨¯©°d·k±kÀ'Ó6Jóôô
 !  t§|¡|Ó4ˆFØ"(ˆD�M‰M˜$ÑØˆMøäò 	VÜ Ð#4°T°FÐ:JÐ!KÓLÐRUÐUûð	Vús   Á A.B/ Â/	CÂ8CÃC)Ú__name__Ú
__module__Ú__qualname__r   r@   rd   Ústrrn   r   r   r   r]   r]   ¥   s/   „ ð3˜vŸ™ð 3°4ó 3ðV ð V¨ô Vr   r]   c                   óœ   — e Zd ZdZdej
                  dej
                  ddfd„Z	 	 	 	 	 ddeeeef   deeeef   d	e	e
   d
ede	e   ddfd„Zy)rf   zT
    Represents a compiled CUDA kernel that can be called with PyTorch tensors.
    rk   r^   r   Nc                 ó    — || _         || _        y r`   )rk   r^   )rc   rk   r^   s      r   rd   z_CudaKernel.__init__Ç   s   € ØˆŒ	Øˆ�r   ÚgridÚblockÚargsÚ
shared_memÚstreamc                 ó–  — ddl }|j                  j                  j                  «       }|sg }g }g }	|D �]O  }
t	        |
|j
                  «      rŒ|
j                  s'|
j                  r|
j                  «       st        d«      ‚t        j                  |
j                  «       «      }|j                  |«       |	j                  t        j                  |«      «       Œ¦t	        |
t        «      r:t        j                   |
«      }|	j                  t        j                  |«      «       Œðt	        |
t"        «      r;t        j$                  |
«      }|	j                  t        j                  |«      «       �Œ;t'        dt)        |
«      › �«      ‚ t        j                  t+        |	«      z  «       }t-        |	«      D ],  \  }}
t        j.                  |
t        j                  «      ||<   Œ. |€ddl}|j                  j3                  «       }t5        |j7                  | j8                  |d   |d   |d   |d   |d   |d   ||j:                  |d«      «       y)aþ  
        Call the compiled CUDA kernel

        Args:
            grid (tuple): Grid dimensions (grid_x, grid_y, grid_z)
            block (tuple): Block dimensions (block_x, block_y, block_z)
            args (list): List of arguments to pass to the kernel.
                         PyTorch tensor arguments will be automatically converted to pointers.
            shared_mem (int): Shared memory size in bytes
            stream (torch.cuda.Stream): CUDA stream to use. If None, uses current stream.
        r   Nz?All tensor arguments must be CUDA tensors or pinned CPU tensorszUnsupported argument type: é   é   )rJ   r7   Ú_utilsr   Ú
isinstanceÚTensorÚis_cudaÚis_cpuÚ	is_pinnedÚ
ValueErrorr   r@   Údata_ptrr<   r   r3   Úc_intÚfloatÚc_floatÚ	TypeErrorÚtyper?   Ú	enumerateÚcastr2   Úcurrent_streamr    ÚcuLaunchKernelrk   Ú_as_parameter_)rc   ru   rv   rw   rx   ry   rJ   r   Úprocessed_argsÚc_argsÚargÚptrr…   r‡   Úc_args_arrayÚis                   r   Ú__call__z_CudaKernel.__call__Ë   sß  € ó& 	à—*‘*×#Ñ#×5Ñ5Ó7ˆáØˆDð 13ˆØˆàó 	KˆCÜ˜#˜uŸ|™|Ô,Ø—{’{¨C¯JªJ¸3¿=¹=¼?Ü$ØYóð ô —o‘o c§l¡l£nÓ5�Ø×%Ñ% cÔ*Ø—‘œfŸl™l¨3Ó/Õ0Ü˜C¤Ô%äŸ™ SÓ)�à—‘œfŸl™l¨5Ó1Õ2ä˜C¤Ô'ä Ÿ.™.¨Ó-�à—‘œfŸl™l¨7Ó3Ö4äÐ"=¼dÀ3»i¸[Ð IÓJÐJð-	Kô2 Ÿ™¬#¨f«+Ñ5Ó8ˆÜ Ó'ò 	@‰FˆAˆsÜ$Ÿk™k¨#¬v¯©Ó?ˆL˜ŠOð	@ð ˆ>ãà—Z‘Z×.Ñ.Ó0ˆFäØ×"Ñ"Ø—	‘	Ø�Q‘Ø�Q‘Ø�Q‘Ø�a‘Ø�a‘Ø�a‘ØØ×%Ñ%ØØóõ	
r   )©r{   r{   r{   r–   Nr   N)ro   rp   rq   Ú__doc__r   r@   rd   Útupler3   r   Úlistr   r•   r   r   r   rf   rf   Â   sž   „ ñð˜VŸ_™_ð °f·o±oð È$ó ð &/Ø&/Ø#ØØ $ñP
à�C˜˜c�MÑ"ðP
ð �S˜#˜s�]Ñ#ðP
ð �t‰nð	P
ð
 ðP
ð ˜‘ðP
ð 
ôP
r   rf   rZ   Úkernel_namesc           
      ó8  — ddl }t        «       }t        | t        «      r| j	                  d«      } t        j                  «       }|j                  j                  «       }|5  t        |j                  t        j                  |«      | «      «       ddd«       |st        |«      S i }|D ]c  }t        j                  «       }t        |j                  t        j                  |«      ||j	                  d«      «      «       t        ||«      ||<   Œe |S # 1 sw Y   Œ‚xY w)a,  
    Loads a CUDA module from PTX code and returns a module object that can access kernels.

    Args:
        ptx (bytes or str): The PTX code to load
        kernel_names (list, optional): List of kernel names to extract from the module.
                                      If None, will return a module object with __getattr__.

    Returns:
        object: If kernel_names is None, returns a module object with __getattr__ to access kernels.
               If kernel_names is provided, returns a dict mapping kernel names to _CudaKernel objects.
    r   Nr0   )r2   r   r~   rr   r6   r   r@   r7   rŒ   r    ÚcuModuleLoadDatar   r]   ri   rf   )	rZ   rš   rJ   r   r^   ry   Úkernelsre   rk   s	            r   Ú_cuda_load_modulerž     sÿ   € ó  ô  Ó!€Gô �#”sÔØ�j‰j˜Ó!ˆô �_‰_Ó€Fà�Z‰Z×&Ñ&Ó(€FØ	ñ IÜ�G×,Ñ,¬V¯\©\¸&Ó-AÀ3ÓGÔH÷Iñ Ü˜6Ó"Ð"ð €GØò 2ˆÜ�‰Ó ˆÜØ×'Ñ'Ü—‘˜TÓ" F¨D¯K©K¸Ó,@óô	
ô
 $ D¨&Ó1ˆ�Šð2ð €N÷!Ið Iús   Á /DÄDÚdeviceÚoptionalÚ	allow_cpuc                 óÐ  — t        | t        «      r| S t        | t        «      rt        j                  | «      } t        | t        j                  «      r;|r| j
                  dvr+t        d| › �«      ‚| j
                  dk7  rt        d| › �«      ‚t        j                  j                  «       s0t        | t        j                  j                  «      r| j                  S t        | ||«      S )a±  Get the device index from :attr:`device`, which can be a torch.device object, a Python integer, or ``None``.

    If :attr:`device` is a torch.device object, returns the device index if it
    is a CUDA device. Note that for a CUDA device without a specified index,
    i.e., ``torch.device('cuda')``, this will return the current default CUDA
    device if :attr:`optional` is ``True``. If :attr:`allow_cpu` is ``True``,
    CPU devices will be accepted and ``-1`` will be returned in this case.

    If :attr:`device` is a Python integer, it is returned as is.

    If :attr:`device` is ``None``, this will return the current default CUDA
    device if :attr:`optional` is ``True``.
    )r7   Úcpuz(Expected a cuda or cpu device, but got: r7   z!Expected a cuda device, but got: )r~   r3   rr   rJ   rŸ   r‰   rƒ   ÚjitÚis_scriptingr7   ÚidxÚ_torch_get_device_index)rŸ   r    r¡   s      r   r   r   N  s·   € ô  �&œ#ÔØˆÜ�&œ#ÔÜ—‘˜fÓ%ˆÜ�&œ%Ÿ,™,Ô'ÙØ�{‰{ /Ñ1Ü Ð#KÈFÈ8Ð!TÓUÐUØ�[‰[˜FÒ"ÜÐ@ÀÀÐIÓJÐJÜ�9‰9×!Ñ!Ô#Ü�fœeŸj™j×/Ñ/Ô0Ø—:‘:ÐÜ" 6¨8°YÓ?Ð?r   )NÚ NNr`   )FF)r   r   Útypingr   r   r   rJ   Útorch._utilsr   r§   r   r   r3   r    r"   rr   r™   Úbytesr[   r]   rf   Údictrž   Úboolr   r   r   ú<module>r®      sK  ðÛ Û 
ß 'Ñ 'ã õ Fð+˜6Ÿ;™;ó +ð	7˜ð 	7 ó 	7ð*˜FŸK™Kó *ð )-ØØ(,Ø#'ñyØðyàðyð ! ™ðyð ð	yð
   ‘~ðyð ˜4‘.ðyð óy÷xVñ V÷:Y
ñ Y
ðz AEñ-Ø	ˆs�EˆzÑ	ð-Ø*2°4¸±9Ñ*=ð-à
ˆ;˜˜S -Ð/Ñ0Ð0Ñ1ó-ðb <Añ@Øð@Øð@Ø48ð@àô@r   