
    #ZjS              
          U d dl Z d dlZd dlZd dlZd dlZd dlZd dlZd dlZd dlm	Z	 d dl
mZ d dlmZmZmZmZmZmZmZmZmZmZmZ d dlZd dlmZmZmZ d dlmZ d dlm Z  dd	l!m"Z"m#Z#m$Z$m%Z%m&Z&m'Z'm(Z(m)Z)m*Z*m+Z+ dd
l!m,Z- ddl!m.Z/ ddl!m0Z0  e	dd          Z1e	e2         e3d<   d dlm4Z4 ddl5m6Z6 ddl6m7Z7m8Z8m9Z9m:Z:m;Z;m<Z< e4rddl=m>Z> didZ?d Z@d ZAdjdeBddfdZCdede2fdZDdede2fdZEdede2fdZFdede2fdZGdede2fd ZHdede2fd!ZIdede2fd"ZJdede2fd#ZKdede2fd$ZLdede2fd%ZMdede2fd&ZNd' ZOd(eBdd)fd*ZPde2fd+ZQdjd(eBde2fd,ZRdjd(eBde2fd-ZSd.e8d/e8ddfd0ZT	 dkd1d2d3e:d4eeB         d5eUde9fd6ZVd7d8d9ee8e<e;ee8         f         d:eBdeBfd;ZWdld<ZXdld=ZYdld>ZZd?eege2f         d@eegef         dedefdAZ[d?eege2f         dedefdBZ\	 	 dmdCe8dDe2dEed         ddFfdGZ]ddHdIdFdJedK         de8fdLZ^	 dndCe8dDe2dMe2ddNfdOZ_ddHdPdNdJedK         de8fdQZ`	 dodCe8dDe2ddRfdSZaddHdTdRdJedK         de8fdUZb edV          ZcdWedXecf         dYedZededXecf         fd[Zd G d\ d]ee          Zfd^egdWeeee2gef         d9ed3eddf
d_Zhe ji        dpda            Zje ji        db             Zke ji        dqdcegddeBfde            Zle G df dg                      Zmg dhZndS )r    N)
ContextVar)	dataclass)AnyCallableDictListMappingOptionalSequenceTupleTypeVarUnioncast)SchemaValidationErrorvalidate_type)Version)table   )
cupycupy_from_dlpackhas_cupyhas_cupy_gpuhas_gpu	has_mxnethas_tensorflow	has_torchhas_torch_cuda_gpuhas_torch_mps)mxnet)
tensorflow)torchDATA_VALIDATIONF)default)TYPE_CHECKING)types)
ArgsKwargsArrayXdFloatsXdIntsXdPaddedRagged)Opsreturntorch.devicec                  `   t           t          d          ddlm}  ddlm} ddlm}  |             }t          ||          r5t           j	        
                                }t          j        d|           S t          ||          rt          j        d          S t          j        d          S )	Nz<Cannot get default Torch device when Torch is not available.r   )get_current_ops)CupyOps)MPSOpszcuda:mpscpu)r"   
ValueErrorbackendsr1   backends.cupy_opsr2   backends.mps_opsr3   
isinstancecudacurrent_devicedevice)r1   r2   r3   ops	device_ids        V/Users/jameslopez/projects/MentorCore/.venv/lib/python3.11/site-packages/thinc/util.pyget_torch_default_devicerA   9   s    }WXXX))))))******((((((
/

C#w #J--//	|/I//000	C	 	  #|E"""<    c                     t          |           rt          S t          |           rt          S t	          dt          |            d          )Nz4Only numpy and cupy arrays are supported, but found z` instead. If get_array_module module wasn't called directly, this might indicate a bug in Thinc.)is_numpy_arraynumpyis_cupy_arrayr   r6   type)arrs    r@   get_array_modulerI   K   s[    c 

	s		 
<99< < <
 
 	
rB   c                      t           S N)r    rB   r@   gpu_is_availablerM   Y   s    NrB   seedc                    t          j        |            t          j                             | dz             t          rt	          j        | dz             t          rzt          j                             |            t          rVt          rQt          j	        
                    |            dt          j        j        _        dt          j        j        _        dS dS dS dS )z@Set the random seed across random, numpy.random and cupy.random.l        l    TFN)randomrN   rE   r   r"   manual_seedr   r   r   r;   manual_seed_allr7   cudnndeterministic	benchmark)rN   s    r@   fix_random_seedrV   ]   s    
K	LdUl### 8 	$!66777 	3 	 	3+ 	3 J&&t,,,15EN .-2EN ***	3 	3	3 	3 	3 	3rB   objc                 >    t          |           pt          |           S )z1Check whether an object is a numpy or cupy array.)rD   rF   rW   s    r@   is_xp_arrayrZ   q   s    #4-"4"44rB   c                 P    t           sdS t          | t          j                  rdS dS )z(Check whether an object is a cupy array.FT)r   r:   r   ndarrayrY   s    r@   rF   rF   v   s-     u	C	&	& turB   c                 >    t          | t          j                  rdS dS )z)Check whether an object is a numpy array.TF)r:   rE   r\   rY   s    r@   rD   rD      s     #u}%% turB   c                 P    t           dS t          | t           j                  rdS dS NFT)r"   r:   TensorrY   s    r@   is_torch_arrayra      s*    }u	C	&	& turB   c                 .    t          |           o| j        S rK   )ra   is_cudarY   s    r@   is_torch_cuda_arrayrd      s    #.3;.rB   c                 >    t          |           pt          |           S rK   )rd   is_torch_mps_arrayrY   s    r@   is_torch_gpu_arrayrg      s    s##>'9#'>'>>rB   c                 N    t          |           ot          | d          o| j        S )Nis_mps)ra   hasattrri   rY   s    r@   rf   rf      s&    #H73#9#9HcjHrB   c                 P    t           sdS t          | t          j                  rdS dS r_   )r   r:   tfr`   rY   s    r@   is_tensorflow_arrayrm      s-     u	C	#	# turB   c                 2    t          |           od| j        v S )NzGPU:)rm   r=   rY   s    r@   is_tensorflow_gpu_arrayro      s    s##<#*(<<rB   c                 Z    t           sdS t          | t          j        j                  rdS dS r_   )r   r:   mxndNDArrayrY   s    r@   is_mxnet_arrayrt      s/     u	C	'	' turB   c                 @    t          |           o| j        j        dk    S )Nr5   )rt   contextdevice_typerY   s    r@   is_mxnet_gpu_arrayrx      s    #C3;#:e#CCrB   c                     t          | t          j                  r| S t          r.t          | t          j                  r|                                 S t          j        |           S rK   )r:   rE   r\   r   r   getarray)datas    r@   to_numpyr}      sT    $&& !	 !jt|44 !xxzz{4   rB   gpu_idzcupy.cuda.Devicec                     t           st          d          t          j        j                            |           }|                                 t          rt          j        	                    |            |S )z=Set the current GPU device for cupy and torch (if available).zNo CUDA GPU devices detected)
r   r6   r   r;   r=   Deviceuser   r"   
set_device)r~   r=   s     r@   set_active_gpur      sa     97888Y$$V,,F
JJLLL &
f%%%MrB   c                  B    ddl m} m}  | d          } ||           dS )z'Use CPU through best available backend.r   )get_opsset_current_opsr5   T)r7   r   r   )r   r   r>   s      r@   require_cpur      s<    22222222
'%..COC4rB   c                 >    t           rt          |            t           S )z?Use GPU if it's available. Returns True if so, False otherwise.r~   )r   require_gpur   s    r@   
prefer_gpur      s      #6""""NrB   c                    ddl m}m}m} t	          j                    dk    r,t          s%t          rt          d          t          d          t	          j                    dk    rt          st          d          t          st          d          t          r# | |                       t          |            n | |                       dS )	Nr   )r2   r3   r   Darwinz6Cannot use GPU, installed PyTorch does not support MPSz(Cannot use GPU, PyTorch is not installedz%Cannot use GPU, CuPy is not installedzNo GPU devices detectedT)r7   r2   r3   r   platformsystemr   r   r6   r   r   r   r   )r~   r2   r3   r   s       r@   r   r      s    ::::::::::H$$]$ 	WUVVVCDDD			h	&	&x	&@AAA 42333 "		"""v!!!4rB   dstsrcc                 "   t          | t          j                  r#t          |t          j                  r	|| d d <   d S t          |           r-t	          j        |d          }t	          j        | |           d S t          j        | |           d S )NF)copy)r:   rE   r\   rF   r   r{   copyto)r   r   s     r@   
copy_arrayr      s    #u}%% *S%-*H*H AAA	s		 j5)))CS#rB           )label_smoothingY	n_classesr   c          	         |$t          t          j        |           dz             }|dk     rt          d          |dk    r|dk    rt          d          d}n!|dk    st          d| d          ||dz
  z  }|dz
  |z  }|dk    r||k    rt          d| d	| d
| d          t	          |           }|                    ||f|d          }|                    |d|z
             ||          S )Nr   r   z>Label-smoothing parameter has to be greater than or equal to 0r   zn_classes should be at least 1zGn_classes should be greater than 1 when label smoothing is enabled,but z was provided.zFor z7 classes label_smoothing parameter has to be less than z, but found .float32)dtype)intrE   maxr6   rI   fullfill_diagonal)r   r   r   nongold_prob
max_smoothxplabel_distrs          r@   to_categoricalr      se    	!q())	L
 
 	
 #>>=>>>1}}1 1 1 1   ')a-8a-9,J1}}J66:9 : :: :'6: : :
 
 	
 
!		B''9i0,i'PPK[!o"5666q>rB   dimXr   c                   t          | t                    rt          | j        |          S t          | t                    rt          | j        |          S t          | d          rt          | d          rxt          t          |           } t          | j	                  dk    rdS t          | j	                  dk    r$t          |                                           dz   S | j	        |         S t          | t          t          f          r,t          |           dk    rdS t          | d         |          S d}t          |          )zrInfer the 'width' of a batch of data, which could be any of: Array,
    Ragged, Padded or Sequence of Arrays.
    r   shapendimr   r   z=Cannot get width of object: has neither shape nor __getitem__)r:   r,   	get_widthr|   r+   rj   r   r(   lenr   r   r   listtupler6   )r   r   errs      r@   r   r   &  s&    !V S))))	Av		 S))))	G		 F!3!3 !qw<<11\\Qquuww<<!##73<	Ae}	%	% q66Q;;1QqTs++++MoorB   c                  ^    d} t           s#t          |                     d                    dS )z4Raise an ImportError if TensorFlow is not installed.z~TensorFlow support requires {pkg}: pip install thinc[tensorflow]

Enable TensorFlow support with thinc.api.enable_tensorflow()ztensorflow>=2.0.0,<2.6.0)pkgN)r   ImportErrorformat)templates    r@   assert_tensorflow_installedr   B  s<     RH K(//.H/IIJJJK KrB   c                  2    t           st          d          dS )z/Raise an ImportError if MXNet is not installed.zjMXNet support requires mxnet: pip install thinc[mxnet]

Enable MXNet support with thinc.api.enable_mxnet()N)r   r   rL   rB   r@   assert_mxnet_installedr   I  s)     
z
 
 	

 
rB   c                  2    t           st          d          dS )z1Raise an ImportError if PyTorch is not installed.z8PyTorch support requires torch: pip install thinc[torch]N)r   r   rL   rB   r@   assert_pytorch_installedr   Q  s&     VTUUUV VrB   is_matchconvert_itemc                 F      |          r |          S t          |t                    rDt           t          |                                                    }t          j        |          S t          |t                    rEi }|                                D ],\  }}t           |          }t           |          }|||<   -|S t          |t                    r fd|D             S t          |t                    rt           fd|D                       S |S )zEither convert a single value if it matches a given function, or
    recursively walk over potentially nested lists, tuples and dicts applying
    the conversion, and returns the same type. Also supports the ArgsKwargs
    dataclass.
    c                 2    g | ]}t          |          S rL   convert_recursive.0itemr   r   s     r@   
<listcomp>z%convert_recursive.<locals>.<listcomp>l  s&    PPPD!(L$??PPPrB   c              3   :   K   | ]}t          |          V  d S rK   r   r   s     r@   	<genexpr>z$convert_recursive.<locals>.<genexpr>n  s0      UU&xtDDUUUUUUrB   )r:   r'   r   r   items
from_itemsdictr   )r   r   rW   	convertedkeyvalues   ``    r@   r   r   W  s<    x}} |C   	C	$	$ %hd399;;>O>OPP	$Y///	C		 	))++ 	# 	#JC#HlC@@C%heDDE"IcNN	C		 PPPPPCPPPP	C		 UUUUUQTUUUUUU
rB   c              #     K    | |          r|V  dS t          |t                    r7t          | t          |                                                    E d{V  dS t          |t
                    rH|                                D ]1\  }}t          | |          E d{V  t          | |          E d{V  2dS t          |t                    st          |t                    r|D ]}t          | |          E d{V  dS dS )zEither yield a single value if it matches a given function, or recursively
    walk over potentially nested lists, tuples and dicts yielding matching
    values. Also supports the ArgsKwargs dataclass.
    N)r:   r'   iterate_recursiver   r   r   r   )r   rW   r   r   r   s        r@   r   r   s  sZ     
 x}} 
9						C	$	$ 9$XtCIIKK/@/@AAAAAAAAAAA	C		 9))++ 	: 	:JC(3777777777(59999999999	: 	: 
C		 9*S%"8"8 9 	9 	9D(488888888889 9	9 	9rB   	xp_tensorrequires_gradr=   ztorch.Tensorc                    t                       |t                      }t          | d          r9|                                 }t          j        j                            |          }nIt          | d          r%t          j        j                            |           }nt	          j        |           }|	                    |          }|r|
                                 |S )z3Convert a numpy or cupy tensor to a PyTorch tensor.NtoDlpack
__dlpack__)r   rA   rj   r   r"   utilsdlpackfrom_dlpack
from_numpytorequires_grad_)r   r   r=   dlpack_tensortorch_tensors        r@   xp2torchr     s     ~)++y*%% 3!**,,{)55mDD	L	)	) 3{)55i@@'	22??6**L &##%%%rB   )r>   r   r>   r-   c                   ddl m} t                       t          |           ryt	          ||          r8|                                                                                                 S t          t          j
        j                            |                     S t	          ||          s|8|                                                                                                 S t          j        |           S )zConvert a torch tensor to a numpy or cupy tensor depending on the `ops` parameter.
    If `ops` is `None`, the type of the resultant tensor will be determined by the source tensor's device.
    r   NumpyOps)apir   r   rd   r:   detachr5   rE   r   r"   r   r   	to_dlpackr   asarray)r   r>   r   s      r@   torch2xpr     s     <(( 	.c8$$ 	P&&((,,..44666#EK$6$@$@$N$NOOOc8$$ 	.&&((,,..44666<---rB   as_variablez	tf.Tensorc                    t                       t          | d          r9|                                 }t          j        j                            |          }n]t          | d          r9|                                 }t          j        j                            |          }nt          j        |           }|rGt          j	        |j	                  5  t          j
        ||          }ddd           n# 1 swxY w Y   |du rI|du rEt          j	        |j	                  5  t          j        |          }ddd           n# 1 swxY w Y   |S )zAConvert a numpy or cupy tensor to a TensorFlow Tensor or Variabler   r   )	trainableNF)r   rj   r   rl   experimentalr   r   r   convert_to_tensorr=   Variablestop_gradient)r   r   r   r   	tf_tensors        r@   xp2tensorflowr     s     !!!y*%% 4!**,,O*66}EE			L	)	) 4!,,..O*66}EE		(33	 H Yy'(( 	H 	HIGGGI	H 	H 	H 	H 	H 	H 	H 	H 	H 	H 	H 	H 	H 	H 	H+"6"6 Yy'(( 	4 	4(33I	4 	4 	4 	4 	4 	4 	4 	4 	4 	4 	4 	4 	4 	4 	4s$   C33C7:C7E  EEr   c                   ddl m} t                       t          |           rWt	          ||          r|                                 S t          j        j        	                    |           }t          |          S t	          ||          s||                                 S t          j        |                                           S )zConvert a Tensorflow tensor to numpy or cupy tensor depending on the `ops` parameter.
    If `ops` is `None`, the type of the resultant tensor will be determined by the source tensor's device.
    r   r   )r   r   r   ro   r:   rE   rl   r   r   r   r   r   r   )r   r>   r   r   s       r@   tensorflow2xpr     s     !!!y)) 
3c8$$ 	3??$$$O2<<YGGM#M222c8$$ 	3??$$$<	 1 1222rB   zmx.nd.NDArrayc                    t                       t          | d          r4|                                 }t          j                            |          }nt          j                            |           }|r|                                 |S )z1Convert a numpy or cupy tensor to a MXNet tensor.r   )r   rj   r   rq   rr   r   r   attach_grad)r   r   r   	mx_tensors       r@   xp2mxnetr     s     y*%% 0!**,,E%%m44		E$$Y//	  rB   r   c                   ddl m} t                       t          |           rWt	          ||          r&|                                                                 S t          |                                           S t	          ||          s|&|                                                                 S t          j
        |                                           S )z1Convert a MXNet tensor to a numpy or cupy tensor.r   r   )r   r   r   rx   r:   r   asnumpyr   to_dlpack_for_writer   r   )r   r>   r   s      r@   mxnet2xpr     s     )$$ 	5c8$$ 	E##%%--///#I$A$A$C$CDDDc8$$ 	5##%%--///<	 1 1 3 3444rB   PartialTfunc.argskwargsc                 H    t          j        | g|R i |}| j        |_        |S )znWrapper around functools.partial that retains docstrings and can include
    other workarounds if needed.
    )	functoolspartial__doc__)r   r   r   partial_funcs       r@   r   r     s4     $T;D;;;F;;L<LrB   c                   v    e Zd Zg fdedededeeeeef                  ee	eef                  f         ddf
dZ
dS )DataValidationErrornamer   r   errorsr.   Nc                    d| d}dt          |           dt          |           }g }|D ]_}d                    d |                    dg           D                       }	|                    |	|                    d          f           `||t	          |          g}
t
                              | d	d
                    |
          z              dS )z8Custom error for validating inputs / outputs at runtime.zData validation error in ''zX: z Y: z -> c                 ,    g | ]}t          |          S rL   )str)r   ps     r@   r   z0DataValidationError.__init__.<locals>.<listcomp>#  s    "H"H"Ha3q66"H"H"HrB   locmsgz


N)rG   joinrz   appendr   r6   __init__)selfr  r   r   r  message	type_infor|   errorerr_locresults              r@   r  zDataValidationError.__init__  s     7t6660$q''00tAww00	 	5 	5Ekk"H"H599UB3G3G"H"H"HIIGKK%))E"2"2344449eDkk2D&499V+<+<"<=====rB   )__name__
__module____qualname__r
  r   r   r   r	   r   r   r  rL   rB   r@   r  r    s         LN> >> > 	>
 hwsCx014S#X3GGH> 
> > > > > >rB   r  r  c                    t          j        |          }t          |j                  }t	          |          dk    rBt	          |           dd                    |           d}d| }t          | ||d|ig          |j        |d                  }|j        |d                  }	g }
|o|j        t          urat          |t                    rt	          |          d
k    r
|d	d
         }t          ||j                  }|r|
                    d|d           |Yt          j        |          j        }|t          j        j        ur-t          |d f|          }|r|
                    d|d           |
rt          | |||
          d	S )zValidate the input and output of a forward function against the type
    annotations, if available. Used in Model.initialize with the input and
    output samples as they pass through the network.
       z (z, )zIInvalid forward function. Expected 3 arguments (model, X, is_train), got r  r      N   )r   )r  r  c                     | S rK   rL   )xs    r@   <lambda>z+validate_fwd_input_output.<locals>.<lambda>B  s    a rB   )r   )r   from_functionr   model_fieldsr   r  r  
annotationr   r:   r   r  inspect	signaturereturn_annotation	Signatureempty)r  r   r   r   schemafields
bad_paramsr   x_fieldy_fieldr  rets               r@   validate_fwd_input_outputr1  )  s    !$''F&%&&F
6{{aF;;tyy'8'8;;;
fZdff!$1s|n===!&),G!&),GF}+366a 	3q66A::"1"AAw122 	7MM&55666}%%7g'---KK 0#66C ;fS99::: 6!$1f5556 6rB   rc              #      K   t          j        | d          }|V  |                                 t          j        |j                   d S )NF)modedelete)tempfileNamedTemporaryFilecloseosremover  )r4  fs     r@   make_tempfiler<  I  sI      #e<<<A
GGGGGIIIIafrB   c              #     K   t          j                    5  t                                          }t                              |            d V  t                              |           d d d            d S # 1 swxY w Y   d S rK   )	threadingLockr#   rz   set)
validationprevs     r@   data_validationrC  Q  s      			 " """$$J'''D!!!	" " " " " " " " " " " " " " " " " "s   AA55A9<A9r  id_colorc              #      K   t           rNt          j        j                            | |           dV  t          j        j                                         dS dV  dS )zxContext manager to register the executed code as an NVTX range. The
    ranges can be used as markers in CUDA profiling.N)r   r   r;   nvtx	RangePushRangePop)r  rD  s     r@   use_nvtx_rangerI  Z  s\        	  (333	!!!!!rB   c                   d    e Zd ZU dZej        ed<   ej        ed<   ede	fd            Z
de	fdZdS )	ArrayInfoz4Container for info for checking array compatibility.r   r   rH   c                 0     | |j         |j                  S )Nr   r   rM  )clsrH   s     r@   
from_arrayzArrayInfo.from_arraym  s    s#)4444rB   c                     |j         | j         k    rt          d| j          d|j                    |j        | j        k    rt          d| j         d|j                   d S )NzShape mismatch in backprop. Y: z, dY: zType mismatch in backprop. Y: )r   r6   r   )r  rH   s     r@   check_consistencyzArrayInfo.check_consistencyq  s|    9
""O$*OOCIOO   9
""NNN39NN   #"rB   N)r  r  r  r  r&   Shape__annotations__DTypesclassmethodr(   rO  rQ  rL   rB   r@   rK  rK  f  sx         >>;<5W 5 5 5 [5W      rB   rK  )rI   rA   rV   rF   rD   r   r   r   r   r   r   r   r   r   r   r1  r  r<  rI  rK  r   r   )r.   r/   )r   rK   )r.   N)FN)FF)F)r2  )r   )o
contextlibr   r&  r9  r   rP   r6  r>  contextvarsr   dataclassesr   typingr   r   r   r   r	   r
   r   r   r   r   r   rE   confection.validationr   r   r   packaging.versionr   wasabir   compatr   r   r   r   r   r   r   r   r   r   r    rq   r!   rl   r"   r#   boolrS  r%    r&   r'   r(   r)   r*   r+   r,   r   r-   rA   rI   rM   r   rV   rZ   rF   rD   ra   rd   rg   rf   rm   ro   rt   rx   r}   r   r   r   r   r   floatr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r6   r  r
  r1  contextmanagerr<  rC  rI  rK  __all__rL   rB   r@   <module>rc     s	             				        " " " " " " ! ! ! ! ! !                           H H H H H H H H H H % % % % % %                                     $ $ $ $ $ $      $.J/@%$P$P$PD! P P P                   H H H H H H H H H H H H H H H H    $
 
 
  3 3# 3d 3 3 3 3(5S 5T 5 5 5 5
s t              /S /T / / / /?C ?D ? ? ? ?IC ID I I I IS T    = = = = = =     DC DD D D D D! ! !3 #5    T     s 4      D    *G ' d      $% !	% % %%}% 	%
 % % % %R IK  Wffhw&778BE   8K K K K
 
 
 
V V V Vud{#3;SE3J3GNQ   89# 5 9C 9C 9 9 9 9(  '+  ^$ 	   8 =A. . . .*25/.. . . ., JO '+BF   6 7;3 3 33$,UO33 3 3 3. /4 '+     ;?5 5 55(055 5 5 5( 7:
3=
!*-9<c8m   > > > > >* > > >&6
6sC.3469<6AD6	6 6 6 6@     " " "  C 3            ,  rB   