+
    &j                        ^ RI t ^ RIt] P                  R R l4       t] P                  R R l4       t] P                  RR R ll4       t] P                  R R	 l4       t] P                  R
 R l4       t] P                  R R l4       t] P                  R R l4       t	] P                  R R l4       t
] P                  R R l4       t] P                  R R l4       tR# )    Nc                $    V ^8  d   QhR\         /#    returnbool)formats   "k/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/utils/_pallas.py__annotate__r      s           c                 4     ^ RI p R#   \         d     R# i ; i)zCheck if JAX is installed.NTF)jaxImportError)r   s    r
   has_jax_packager      s     s    c                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r      s      D r   c                 \    \        4       '       g   R#  ^ RIHp  R#   \         d     R# i ; i)z0Check if Pallas (JAX experimental) is available.F)pallasT)r   jax.experimentalr   r   )pls    r
   has_pallas_packager      s1     	
  s    ++c                    V ^8  d   QhR\         \        \        \        3,          R\         \        \        \        3,          /# )r   fallbackr   )tupleint)r	   s   "r
   r   r   !   s1     	 	eCcM2 	5cSVCW 	r   c                     ^ RI pVP                  P                  R4      pR VR,           4       w  r4pW4V3#   \        \        \
        3 d    T u # i ; i)z/Get JAX version as (major, minor, patch) tuple.N.c              3   8   "   T F  p\        V4      x  K  	  R # 5i)N)r   ).0vs   & r
   	<genexpr>"get_jax_version.<locals>.<genexpr>'   s     A/@!s1vv/@s   :N   N)r   __version__splitr   
ValueErrorAttributeError)r   r   version_partsmajorminorpatchs   &     r
   get_jax_versionr+       sX    --c2A}R/@Aee$$^4 s   58 AAc                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   .   s      d r   c                 8   \        4       '       g   R#  ^ RIp V P                  R4      p\        V4      ^ 8X  d   R# \        P
                  P                  4       '       d*   \        P
                  P                  4       w  r#V^	8  d   R# R#   \         d     R# i ; i)zJCheck if JAX has CUDA backend support with SM90+ (required by Mosaic GPU).FNgpuT)	r   r   deviceslentorchcudais_availableget_device_capability	Exception)r   r/   r(   r)   s       r
   has_jax_cuda_backendr6   -   s      ++e$w<1 ::""$$ ::;;=LEqy s   %B
 #B
 'B
 
BBc                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   F   s      T r   c                     \        4       '       g   R#  ^ RIp V P                  R4      p\        V4      ^ 8  #   \         d     R# i ; i)z%Check if JAX has TPU backend support.FNtpu)r   r   r/   r0   r5   )r   r/   s     r
   has_jax_tpu_backendr:   E   sI      ++e$7|a s   "7 AAc                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   U   s     	 	t 	r   c                 t     ^ RI p V P                  P                  4        R#   \        \        3 d     R# i ; i)z.Check if torch_tpu is installed and available.NTF)torch_tpu.apiapi
tpu_devicer   RuntimeError)	torch_tpus    r
   has_torch_tpurB   T   s6     	  "& s   " 77c                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   b   s          r   c                     \        4       # )z,Checks for a full Pallas-on-CPU environment.)r    r   r
   has_cpu_pallasrF   a   s     r   c                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   h   s     Y Y Yr   c                     \        4       ;'       d1    \        P                  P                  4       ;'       d    \	        4       # )z-Checks for a full Pallas-on-CUDA environment.)r   r1   r2   r3   r6   rE   r   r
   has_cuda_pallasrI   g   s.     XXEJJ$;$;$=XXBVBXXr   c                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   n   s     N N Nr   c                 ^    \        4       ;'       d    \        4       ;'       d    \        4       # )z,Checks for a full Pallas-on-TPU environment.)r   r:   rB   rE   r   r
   has_tpu_pallasrL   m   s#     MM$7$9MMmoMr   c                $    V ^8  d   QhR\         /# r   r   )r	   s   "r
   r   r   t   s     	E 	ED 	Er   c                 ^    \        4       ;'       g    \        4       ;'       g    \        4       # )z
Check if Pallas backend is fully available for use.

Requirements:
- JAX package installed
- Pallas (jax.experimental.pallas) available
- A compatible backend (CUDA or TPU) is available in both PyTorch and JAX.
)rF   rI   rL   rE   r   r
   
has_pallasrO   s   s#     DD0DDN4DDr   ))r   r   r   )	functoolsr1   cacher   r   r+   r6   r:   rB   rF   rI   rL   rO   rE   r   r
   <module>rR      s         	 	  .   	 	    
 Y Y
 N N
 	E 	Er   