+
    &jR                     .   ^ RI t ^ RIt^ RIHt ^ RIt ^ RIHt Rt^ RI
Ht R R ltR	 R
 ltR R ltR R ltR R ltR R ltR R ltR R ltR R ltR#R R llt ! R R4      t ! R R4      tR$R R  lltR%R! R" lltR#   ]	 d    RtRt Lwi ; i)&    NAny)runtimeTF)_get_device_indexc                8    V ^8  d   QhR\         P                  /#    returnctypesCDLL)formats   "i/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/cuda/_utils.py__annotate__r      s      &++     c                      ^ RI p \        P                  ! \        V P	                  R4      ^ ,          4      4      pVP                  Vn        VP                  Vn        VP                   Vn        VP$                  Vn        VP(                  Vn        V#   \
        \        3 dj    \        P                  R8X  d<   \        P                  ! R\        P                  P                  ^ ,           R24      p L\        P                  ! R4      p Li ; i)r   Namdhip64win32	amdhip64_.dllzlibamdhip64.so)rocm_sdkr   r   strfind_librariesImportError
IndexErrorsysplatformtorchversionhiphipGetErrorStringcuGetErrorStringhipModuleLoadDatacuModuleLoadDatahipModuleGetFunctioncuModuleGetFunctionhipModuleLaunchKernelcuLaunchKernelhipFuncSetAttributecuFuncSetAttribute)r   libs     r   _get_hip_runtime_libraryr,      s    	0kk#h55jA!DEF 00C00C!66C22C 44CJ $ 0<<7"++	%--*;*;A*>)?tDEC++./C	0s   9B AD3DDc                8    V ^8  d   QhR\         P                  /# r   r   )r   s   "r   r   r   -   s     + +6;; +r   c                      \         P                  R 8X  d   \        P                  ! R4      # \        P                  ! R4      # )r   z
nvcuda.dllzlibcuda.so.1)r   r   r   r    r   r   _get_cuda_libraryr0   -   s,    
||w{{<(({{>**r   c                8    V ^8  d   QhR\         P                  /# r   r   )r   s   "r   r   r   5   s     # #&++ #r   c                  j    \         P                  P                  '       d   \        4       # \	        4       # N)r   r   r    r,   r0   r/   r   r   _get_gpu_runtime_libraryr4   5   s$    }}')) ""r   c                (    V ^8  d   QhR\         RR/# r	   resultr
   Nint)r   s   "r   r   r   =   s     	7 	7 	7 	7r   c                    V ^ 8X  d   R# \         P                  ! 4       p\        4       pVP                  V \         P                  ! V4      4       VP
                  e   VP
                  P                  4       MRp\        RV 24      h)r   NUnknown CUDA errorCUDA error: )r   c_char_pr4   r"   byrefvaluedecodeRuntimeError)r7   err_strlibcudaerror_messages   &   r   _check_cudarE   =   sn    {ooG&(GVV\\'%:;")--";AU  m_5
66r   c                0    V ^8  d   QhR\         R\         /# )r	   r7   r
   r   )r   s   "r   r   r   I   s        r   c                n   \         '       g   \        R4      hV vrV\        P                  P                  8w  dQ   \        P
                  ! V4      w  r4\        V\        4      '       d   VP                  4       p\        RV RV R24      h\        V4      ^ 8X  d   R# \        V4      ^8X  d
   V^ ,          # V# )a  Check a cuda.bindings (cuda-python) call result for errors.

All cuda.bindings runtime calls return ``(error, *outputs)``.  This
helper unpacks the tuple, raises on non-success, and returns the
outputs (``None`` for zero outputs, scalar for one, tuple otherwise).
zcuda.bindings is not availabler<   z ()N)
_HAS_CUDA_BINDINGSrA   _cuda_bindings_runtimecudaError_tcudaSuccesscudaGetErrorString
isinstancebytesr@   len)r7   errout_rB   s   &    r   _check_cuda_bindingsrT   I   s     ;<<IC!--99	: #55 	
 gu%%nn&G\#b	;<<
3x1}
3x1}1vJr   c                8    V ^8  d   QhR\         P                  /# r   r   )r   s   "r   r   r   f   s      V[[ r   c                  4    ^ RI p \        P                  ! \        V P	                  R4      ^ ,          4      4      pVP                  Vn        VP                  Vn        VP"                  Vn        VP&                  Vn        VP*                  Vn        VP.                  Vn        VP2                  Vn        VP6                  Vn        VP:                  Vn        VP>                  Vn         V#   \
        \        3 d    \        P                  R8X  dq   RP                  R\        P                  P                  ^ ,          R\        P                  P                  ^,          .4      p\        P                  ! RT R24      p ELB\        P                  ! R4      p EL[i ; i)r   Nhiprtcr    0r   zlibhiprtc.so)!r   r   r   r   r   r   r   r   r   joinr   r   r    hiprtcGetErrorStringnvrtcGetErrorStringhiprtcCreateProgramnvrtcCreateProgramhiprtcDestroyProgramnvrtcDestroyProgramhiprtcCompileProgramnvrtcCompileProgramhiprtcGetCodeSizenvrtcGetCUBINSizehiprtcGetCodenvrtcGetCUBINhiprtcGetProgramLogSizenvrtcGetProgramLogSizehiprtcGetProgramLognvrtcGetProgramLoghiprtcAddNameExpressionnvrtcAddNameExpressionhiprtcGetLoweredNamenvrtcGetLoweredName)r   r+   version_strs      r   _get_hiprtc_libraryrp   f   sA   .kk#h55h?BCD "66C 44C!66C!66C11C))C!$!<!<C 44C!$!<!<C!66CJ) $ .<<7"''emm''*C1B1B11EFK ++{m489C++n-C.s   9C' 'BF=FFc                8    V ^8  d   QhR\         P                  /# r   r   )r   s   "r   r   r      s     6 6FKK 6r   c                  6   \        \        P                  P                  P	                  R 4      ^ ,          4      p \
        P                  R8X  d	   RV  R2.pMRV  2R.pV F  p \        P                  ! V4      u # 	  \        R4      h  \         d     K7  i ; i).r   nvrtc64_z0_0.dllzlibnvrtc.so.zlibnvrtc.soz Could not find any NVRTC library)
r9   r   r   cudasplitr   r   r   r   OSError)major_version
nvrtc_libslib_names      r   _get_nvrtc_libraryr{      s    **005a89M
||w}oW-


 =/*

 	;;x(( 
 4
55  		s   $B		BBc                8    V ^8  d   QhR\         P                  /# r   r   )r   s   "r   r   r      s     $ $fkk $r   c                  j    \         P                  P                  '       d   \        4       # \	        4       # r3   )r   r   r    rp   r{   r/   r   r   _get_gpu_rtc_libraryr~      s&     }}"$$!##r   c                :    V ^8  d   QhR\         \        ,          /# r   )listr   )r   s   "r   r   r      s      tCy r   c                     ^ RI Hp Hp R0pV Uu. uF  q3V9  g   K  VNK  	  pp\        P                  P
                  '       d   VP                  V 4       V# u upi )z
Get HIPCC/NVCC flags that are compatible with NVRTC compilation.

Returns:
    List of HIPCC/NVCC flags that can be safely used with NVRTC.
)COMMON_HIPCC_FLAGSCOMMON_NVCC_FLAGSz--expt-relaxed-constexpr)torch.utils.cpp_extensionr   r   r   r   r    extend)r   r   nvrtc_unsupported_flagsflagcompatible_flagss        r   _get_gpu_rtc_compatible_flagsr      sh     P 	# +*:Q.Q*   }} 23s
   AAc                    V ^8  d   QhR\         R\         R\         R,          R\        R,          R\        R,          R\        R\        \        \         3,          /# )	r	   kernel_sourcekernel_namecompute_capabilityNcuda_include_dirsnvcc_optionsauto_pchr
   )r   r   booltuplerO   )r   s   "r   r   r      sl     O$ O$O$O$ d
O$ d{	O$
 +O$ O$ 5#:O$r   c           
     p	  aa ^ RI p\        4       o^ oR VV3R llpV P                  R4      pVfx   VP                  P	                  VP                  P                  4       4      p	VP                  P                  '       d   V	P                   pMV	P                   V	P                   2p. p
VP                  P                  '       d$   V
P                  RV 2P                  4       4       M"V
P                  RV 2P                  4       4       ^ RIHp V! R4      pV F%  pV
P                  R	V 2P                  4       4       K'  	  V'       d,   V F%  pV
P                  R	V 2P                  4       4       K'  	  V'       d^   \        VP                  P                  4      R
8  d#   \        RVP                  P                   24      hVf   . pVP                  R4       V'       d*   V F#  pV
P                  VP                  R4      4       K%  	  \!        4       pT
P#                  V Uu. uF  pVP                  R4      NK  	  up4       \%        V
4      p\&        P(                  V,          ! V
!  p\&        P*                  ! 4       pV! SP-                  \&        P.                  ! V4      VV R2P                  4       ^ RR4      4       VP                  R4      pV! SP1                  VV4      4       SP3                  VVV4      pVS8w  d   \&        P4                  ! 4       pSP7                  V\&        P.                  ! V4      4       \&        P8                  ! VP:                  4      pSP=                  VV4       \?        RVP:                  PA                  4        24      h\&        P4                  ! 4       pV! SPC                  V\&        P.                  ! V4      4      4       \&        P8                  ! VP:                  4      pV! SPE                  VV4      4       \&        P(                  ! 4       pV! SPG                  VV\&        P.                  ! V4      4      4       VP:                  e   VP:                  PA                  4       pMRpSPI                  \&        P.                  ! V4      4       VPJ                  V3# u upi )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.
    cuda_include_dirs (list, None): List of directories containing CUDA headers
    nvcc_options (list, None): Additional options to pass to NVRTC
    auto_pch (bool): Enable automatic precompiled headers (CUDA 12.8+)

Returns:
    Tuple[bytes, str]: The compiled PTX code and mangled kernel name
Nc                (    V ^8  d   QhR\         RR/# r6   r8   )r   s   "r   r   $_nvrtc_compile.<locals>.__annotate__   s     	? 	?C 	?D 	?r   c                    < V S8w  dt   \         P                  ! 4       pSP                  V \         P                  ! V4      4       VP                  e   VP                  P                  4       MRp\        RV 24      hR # )Nr;   r<   )r   r=   r\   r>   r?   r@   rA   )r7   rB   rD   NVRTC_SUCCESSlibnvrtcs   &  r   check_nvrtc#_nvrtc_compile.<locals>.check_nvrtc   so    ]"oo'G((g1FG ==, $$&) 
 m_=>> #r   utf-8z--offload-arch=z--gpu-architecture=sm_)include_pathsru   z-Iz12.8zPCH requires CUDA 12.8+, got z--pchz.cuzKernel compilation failed:
rX   )&
torch.cudar~   encoderu   get_device_propertiescurrent_devicer   r    gcnArchNamemajorminorappendr   r   r   AssertionErrorr   r   rP   r   r=   c_void_pr^   r>   rl   rb   c_size_trh   create_string_bufferr?   rj   rA   r@   rd   rf   rn   r`   raw)r   r   r   r   r   r   r   r   source_bytespropsoptionsr   cuda_include_paths	cuda_path	directoryoptionnvrtc_compatible_flagsr   num_optionsoptions_arrayprogc_kernel_namereslog_sizelogbinary_sizebinaryc_mangled_namemangled_namer   r   s   &&&&&&                       @@r   _nvrtc_compiler      s   0  $%H M	? 	? !''0L !

001J1J1LM==$)$5$5#6$)KK=!> G}});(<=DDFG/0B/CDKKMN 8&v.'	I;'..01 ( *INNR	{+2245 + u}}!!"V+ #@ASAS@T!UVVLG$ "FNN6==12 # ;<NN5KL5KTDKK(5KLM g,K__{2W=M ??D##LLm3&&(	
	  &&w/M//mDE 
&
&t[-
HC m??$''fll8.DE))(..9##D#.9#)):J:J:L9MNOO //#K**4k1JKL(():):;F&&tV45 __&N$$T=&,,~:VW '%++224  d!34 ::|##o Ms   'R3c                   D   a  ] tR tRt o V 3R lR ltV 3R lR ltRtV tR# )_CudaModuleiI  c                8   < V ^8  d   QhRS[ P                  RR/# )r	   moduler
   Nr   r   )r   __classdict__s   "r   r   _CudaModule.__annotate__J  s     3 3v 34 3r   c                     Wn         / V n        R # r3   )_module_kernels)selfr   s   &&r   __init___CudaModule.__init__J  s    02r   c                $   < V ^8  d   QhRS[ RR/# )r	   namer
   _CudaKernel)r   )r   r   s   "r   r   r   N  s     V V V Vr   c           	        WP                   9   d   V P                   V,          # ^ RIHp V! 4       p\        P                  ! 4       p \        VP                  \        P                  ! V4      V P                  VP                  R4      4      4       \        W@P                  4      pWPP                   V&   V#   \         d   p\        RT R24      ThRp?ii ; i)r   )r4   r   zNo kernel named 'z' in this moduleN)r   torch.cuda._utilsr4   r   r   rE   r&   r>   r   r   r   rA   AttributeError)r   r   r4   rC   funckernelrQ   s   &&     r   __getattr___CudaModule.__getattr__N  s    == ==&& 	?*, 	V++LL&dkk'6J
 !||4F"(MM$M 	V #4TF:J!KLRUU	Vs   A-B5 5C CC)r   r   N)__name__
__module____qualname____firstlineno__r   r   __static_attributes____classdictcell__r   s   @r   r   r   I  s     3 3V Vr   r   c                   ^   a  ] tR tRt o RtV 3R lR ltRV 3R lR lltV 3R lR	 ltR
tV t	R# )r   ig  zL
Represents a compiled CUDA kernel that can be called with PyTorch tensors.
c                R   < V ^8  d   QhRS[ P                  RS[ P                  RR/# )r	   r   r   r
   Nr   )r   r   s   "r   r   _CudaKernel.__annotate__l  s*     ' 'V__ 'foo '$ 'r   c                ,    Wn         W n        ^ V n        R# )r   N)r   r   _max_shared_mem_bytes)r   r   r   s   &&&r   r   _CudaKernel.__init__l  s    	%&"r   Nc                   < V ^8  d   QhRS[ S[S[S[3,          RS[ S[S[S[3,          RS[R,          RS[RS[R,          RR/# )r	   gridblockargsN
shared_memstreamr
   )r   r9   r   r   )r   r   s   "r   r   r   q  sm     _
 _
CcM"_
 S#s]#_
 Tk	_

 _
 d
_
 
_
r   c                   ^ RI pVP                  P                  P                  4       pV'       g   . p. p. p	V EF|  p
\	        WP
                  4      '       d   V
P                  '       g4   V
P                  '       d   V
P                  4       '       g   \        R4      h\        P                  ! V
P                  4       4      pVP                  V4       V	P                  \        P                  ! V4      4       K  \	        V
\        4      '       d?   \        P                   ! V
4      pV	P                  \        P                  ! V4      4       EK  \	        V
\"        4      '       d?   \        P$                  ! V
4      pV	P                  \        P                  ! V4      4       EKh  \'        R\)        V
4       24      h	  \        P                  \+        V	4      ,          ! 4       p\-        V	4       F,  w  r\        P.                  ! V
\        P                  4      W&   K.  	  Vf   ^ RIpVP                  P3                  4       pVR
8  dW   V P4                  ^ 8X  g   W@P4                  8  d6   V P4                  ^ 8X  d   RMRV P4                   R2p\7        RV RV R	24      h\9        VP;                  V P<                  V^ ,          V^,          V^,          V^ ,          V^,          V^,          VVP>                  VR4      4       R# )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.
Nz?All tensor arguments must be CUDA tensors or pinned CPU tensorszUnsupported argument type: znot configuredzonly z bytes configuredzKernel requires z' bytes of shared memory (>= 48KB), but ze. Call kernel.set_shared_memory_config(shared_mem) after compilation and before launching the kernel.   ) r   ru   _utilsr4   rN   Tensoris_cudais_cpu	is_pinned
ValueErrorr   r   data_ptrr   r>   r9   c_intfloatc_double	TypeErrortyperP   	enumeratecastr   current_streamr   rA   rE   r(   r   _as_parameter_)r   r   r   r   r   r   r   rC   processed_argsc_argsargptrr   r   c_args_arrayiconfigured_msgs   &&&&&&           r   __call___CudaKernel.__call__q  sX   & 	**##<<>D 13C#||,,{{{CJJJ3==??$Y  ooclln5%%c*fll3/0C%%S)fll512C''!??3/fll845"=d3i[ IJJ+ 0 #f+58'FA$kk#v?LO ( >ZZ..0F "&&!+z<V<V/V --2 !T7788IJ 
 ":, /%& '33  	""		QQQaaa%%	
r   c                $   < V ^8  d   QhRS[ RR/# )r	   shared_mem_bytesr
   Nr8   )r   r   s   "r   r   r     s     (6 (6 (6 (6r   c                   VR8  d	   Wn         R# \        4       p\        P                  P	                  4       p\        P
                  P                  '       d   VP                  R8w  d   RMR	pM\        VRR4      pW8  d   \        RV RV R24      h^p\        VP                  V P                  VV4      4       Wn         R# )
0   Ngfx950i   shared_memory_per_block_optinr   zRequested shared memory (z bytes) exceeds device limit (z= bytes). Consider reducing block size or shared memory usage.i  )r   r4   r   ru   r   r   r    r   getattrrA   rE   r*   r   )r   r  rC   device_propsmax_shared_mem+cudaFuncAttributeMaxDynamicSharedMemorySizes   &&    r   set_shared_memory_config$_CudaKernel.set_shared_memory_config  s    i')9&*, zz779== &11X=:  %=uN ,+,<+= >!!/ 0 1GG  783&&		; 	
 &6"r   )r   r   r   )   r  r  r  Nr   N)
r   r   r   r   __doc__r   r  r  r   r   r   s   @r   r   r   g  s+     ' '
_
 _
B(6 (6r   r   c          	          V ^8  d   QhR\         \        ,          R\        \         ,          R,          R\        \        \         R3,          ,          /# )r	   ptxkernel_namesNr
   r   )r   rO   r   r   dict)r   s   "r   r   r     s@     - -	u-$(I$4-4]*++-r   c           
     t   ^ RI p\        4       p\        V \        4      '       d   V P	                  R4      p \
        P                  ! 4       pVP                  P                  4       pV;_uu_ 4        \        VP                  \
        P                  ! V4      V 4      4       RRR4       V'       g   \        V4      # / pV Fc  p\
        P                  ! 4       p\        VP                  \
        P                  ! V4      WGP	                  R4      4      4       \        W4      Wg&   Ke  	  V#   + '       g   i     L; i)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.
Nr   )r   r4   rN   r   r   r   r   ru   r   rE   r$   r>   r   r&   r   )	r  r  r   rC   r   r   kernelsr   r   s	   &&       r   _cuda_load_moduler    s       '(G #sjj! __FZZ&&(F	G,,V\\&-A3GH 
 6"" G ''T"FKK,@	

 $D1  N! 
s   /0D''D7	c                H    V ^8  d   QhR\         R\        R\        R\        /# )r	   deviceoptional	allow_cpur
   )r   r   r9   )r   s   "r   r   r   -  s2     @ @@@48@@r   c                $   \        V \        4      '       d   V # \        V \        4      '       d   \        P                  ! V 4      p \        V \        P                  4      '       dH   V'       d!   V P
                  R9  d   \        RV  24      hMV P
                  R8w  d   \        RV  24      h\        P                  P                  4       '       g7   \        V \        P                  P                  4      '       d   V P                  # \        WV4      # )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``.
ru   z(Expected a cuda or cpu device, but got: z!Expected a cuda device, but got: )ru   cpu)rN   r9   r   r   r  r   r   jitis_scriptingru   idx_torch_get_device_index)r  r  r  s   &&&r   r   r   -  s      &#&#f%&%,,''{{/1 #KF8!TUU 2[[F"@IJJ99!!##fejj//00::"6Y??r   )NNNFr3   )FF)r   r   typingr   r   cuda.bindingsr   rJ   rI   r   torch._utilsr   r"  r,   r0   r4   rE   rT   rp   r{   r~   r   r   r   r   r  r/   r   r   <module>r&     s     
    F.+#	7::6&$0O$dV V<S6 S6l-`@ @  !s   B BB