ó
    Eñi(,  ã                  óÂ   • S SK Jr  S SKJrJr  S SKJrJrJrJ	r	  S SK
r
S SKr
S SKJr  S SKJr  SSKJrJrJr  \(       a
  S S	KJr  S SKr " S
 S\5      r " S S\5      rg)é    )Úannotations)ÚABCÚabstractmethod)ÚAnyÚOptionalÚTYPE_CHECKINGÚUnionN)Úir)ÚVé   )ÚFixedLayoutÚFlexibleLayoutÚLayout)ÚSequencec                  óô   • \ rS rSrSr  S     SS jjrSSS jjr\SS j5       r\SS j5       r	SS jr
SS	 jrSS
 jrSS jrSS jrSS jrSS jrSS S jjr\S!S j5       rS"S jr\S#S$S jj5       rSrg)%ÚKernelInputsé   zÎ
Class to store and provide access to input nodes for kernels.
This class takes in a tuple of input nodes and provides methods to access
information about these nodes, such as their device type and device.
Nc                ón   • Xl         SU l        Ub  UO0 U l        X0l        [	        U5      S:”  d   S5       eg)z�
Initialize with a tuple of input nodes.

Args:
    input_nodes: A tuple of input nodes to store
    out_dtype: Optional output dtype to store
Nr   z Expected at least one input node)Ú_input_nodesÚ_device_nameÚ_scalarsÚ
_out_dtypeÚlen)ÚselfÚinput_nodesÚscalarsÚ	out_dtypes       ÚZ/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/_inductor/kernel_inputs.pyÚ__init__ÚKernelInputs.__init__   s>   € ð (ÔØ+/ˆÔØ#*Ñ#6™¸BˆŒØ#ŒÜ�;Ó !Ó#ÐGÐ%GÓGÑ#ó    c                óþ   • Uc  U R                   $ [        U R                   5      [        U5      :X  d)   S[        U R                   5       S[        U5       35       eU Vs/ s H  o R                   U   PM     sn$ s  snf )zÿ
Return the stored input nodes, optionally reordered.

Args:
    reorder: Optional sequence of indices to reorder the nodes.
            For example, (2, 0, 1) would return nodes in that order.

Returns:
    The tuple of input nodes, optionally reordered
zReorder length mismatch: ú vs )r   r   )r   ÚreorderÚis      r   ÚnodesÚKernelInputs.nodes.   s|   € ð ‰?Ø×$Ñ$Ð$Ü�4×$Ñ$Ó%¬¨W«Ó5ð 	
Ø'¬¨D×,=Ñ,=Ó(>Ð'?¸tÄCÈÃLÀ>ÐRó	
Ð5ñ /6Ó6ªg¨×!Ñ! !Ô$©gÑ6Ð6ùÒ6s   ÁA:c                ó,   • [        U R                  5      $ )zH
Get the number of input nodes.

Returns:
    The number of input nodes
)r   r   ©r   s    r   ÚcountÚKernelInputs.count@   s   € ô �4×$Ñ$Ó%Ð%r!   c                óH   • [         R                  " U R                  S   5      $ )z\
Get the device type of the first node.

Returns:
    The device type (e.g., 'cuda', 'cpu')
r   )r
   Úget_device_typer   r)   s    r   Údevice_typeÚKernelInputs.device_typeJ   s    € ô ×!Ò! $×"3Ñ"3°AÑ"6Ó7Ð7r!   c                ó<   • U R                   S   R                  5       $ )zN
Get the device of the first node.

Returns:
    The device of the first node
r   )r   Ú
get_devicer)   s    r   ÚdeviceÚKernelInputs.deviceU   s   € ð × Ñ  Ñ#×.Ñ.Ó0Ð0r!   c                óÔ   • U R                   cP  U R                  5       nU R                  S:X  a0  [        R                  R                  U5      nUR                  U l         U R                   $ )zU
Get the device name information.

Returns:
    A tuple of (gpu_name, vendor, model)
Úcuda)r   r2   r.   Útorchr5   Úget_device_propertiesÚgcnArchName)r   r2   Údevice_propertiess      r   Údevice_nameÚKernelInputs.device_name^   sX   € ð ×ÑÑ$Ø—[‘[“]ˆFØ×Ñ 6Ó)Ü$)§J¡J×$DÑ$DÀVÓ$LÐ!Ø$5×$AÑ$A�Ô!Ø× Ñ Ð r!   c                ó:   • [        S U R                   5       5      $ )zg
Get the symbolic shapes of all input nodes.

Returns:
    A tuple of shape tuples for each input node
c              3  ó@   #   • U  H  oR                  5       v •  M     g 7f©N)Úget_size©Ú.0Únodes     r   Ú	<genexpr>Ú/KernelInputs.shapes_symbolic.<locals>.<genexpr>s   s   é € ÐCÒ1B¨—]‘]—_�_Ò1Bùó   ‚©Útupler   r)   s    r   Úshapes_symbolicÚKernelInputs.shapes_symbolicl   s   € ô ÑC°×1BÒ1BÓCÓCÐCr!   c                ó:   • [        S U R                   5       5      $ )z€
Get the size hints for shapes of all input nodes.

Returns:
    A tuple of shape tuples with integer hints for each input node
c              3  óÒ   #   • U  H]  n[         R                  R                  R                  UR	                  5       [
        R                  R                  R                  S 9v •  M_     g7f©)ÚfallbackN)	r   ÚgraphÚsizevarsÚ
size_hintsr?   r6   Ú	_inductorÚconfigÚunbacked_symint_fallbackr@   s     r   rC   Ú-KernelInputs.shapes_hinted.<locals>.<genexpr>|   sQ   é € ð 
ò
 *�ô	 �G‰G×Ñ×'Ñ'Ø—‘“ÜŸ™×/Ñ/×HÑHð (õ ò *ùó   ‚A%A'rF   r)   s    r   Úshapes_hintedÚKernelInputs.shapes_hintedu   ó&   € ô ñ 
ð
 ×)Ò)ó
ó 
ð 	
r!   c                ó:   • [        S U R                   5       5      $ )zi
Get the symbolic strides of all input nodes.

Returns:
    A tuple of stride tuples for each input node
c              3  ó@   #   • U  H  oR                  5       v •  M     g 7fr>   )Ú
get_strider@   s     r   rC   Ú0KernelInputs.strides_symbolic.<locals>.<genexpr>‹   s   é € ÐEÒ3D¨4—_‘_×&Ð&Ò3DùrE   rF   r)   s    r   Ústrides_symbolicÚKernelInputs.strides_symbolic„   s   € ô ÑE°4×3DÒ3DÓEÓEÐEr!   c                ó:   • [        S U R                   5       5      $ )z‚
Get the size hints for strides of all input nodes.

Returns:
    A tuple of stride tuples with integer hints for each input node
c              3  óÒ   #   • U  H]  n[         R                  R                  R                  UR	                  5       [
        R                  R                  R                  S 9v •  M_     g7frL   )	r   rN   rO   rP   r[   r6   rQ   rR   rS   r@   s     r   rC   Ú.KernelInputs.strides_hinted.<locals>.<genexpr>”   sR   é € ð 
ò
 *�ô	 �G‰G×Ñ×'Ñ'Ø—‘Ó!ÜŸ™×/Ñ/×HÑHð (õ ò *ùrU   rF   r)   s    r   Ústrides_hintedÚKernelInputs.strides_hinted�   rX   r!   c                ó:   • [        S U R                   5       5      $ )zX
Get the dtypes of all input nodes.

Returns:
    A tuple of dtypes for each input node
c              3  ó@   #   • U  H  oR                  5       v •  M     g 7fr>   )Ú	get_dtyper@   s     r   rC   Ú&KernelInputs.dtypes.<locals>.<genexpr>£   s   é € ÐDÒ2C¨$—^‘^×%Ð%Ò2CùrE   rF   r)   s    r   ÚdtypesÚKernelInputs.dtypesœ   s   € ô ÑD°$×2CÒ2CÓDÓDÐDr!   c                ó<   • U R                   U   R                  5       $ )z¢
Get the dtype of a specific input node.

Args:
    idx: Index of the node to get the dtype from (default: 0)

Returns:
    The dtype of the specified input node
)r   rf   )r   Úidxs     r   ÚdtypeÚKernelInputs.dtype¥   s   € ð × Ñ  Ñ%×/Ñ/Ó1Ð1r!   c                ó   • g)úc
Get the output dtype, whether passed in or inferred from the nodes

Returns:
    The output dtype
N© r)   s    r   r   ÚKernelInputs.out_dtype±   ó   � r!   c                óT   • XR                   ;   d   SU S35       eU R                   U   $ )zr
Get the scalar value for a given name.

Args:
    name: Name of the scalar to get

Returns:
    The scalar value
zScalar z not found, but required)r   )r   Únames     r   Ú
get_scalarÚKernelInputs.get_scalarº   s2   € ð —}‘}Ó$ÐN¨°¨vÐ5MÐ&NÓNÐ$Ø�}‰}˜TÑ"Ð"r!   c                ó   • g)zÉ
Abstract method to handle output layout generation.

Args:
    out_dtype: Optional output dtype. If not provided, infer from inputs
    flexible: If True, return FlexibleLayout, otherwise FixedLayout
Nrp   )r   Úflexibles     r   Úoutput_layoutÚKernelInputs.output_layoutÇ   rr   r!   )r   r   r   r   )NN)r   ú	list[Any]r   ú&Optional[dict[str, Union[float, int]]]r   úOptional[torch.dtype]r>   )r$   zOptional[Sequence[int]]Úreturnr{   ©r~   Úint)r~   zOptional[str])r~   ztorch.device)r~   ztuple[tuple[Any, ...], ...])r~   ztuple[tuple[int, ...], ...])r~   z%tuple[tuple[sympy.Integer, ...], ...])r~   ztuple[torch.dtype, ...])r   )rk   r€   r~   útorch.dtype©r~   r�   )rt   Ústrr~   zUnion[float, int]©T©rx   Úboolr~   r   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   r&   Úpropertyr*   r.   r2   r:   rH   rV   r]   rb   rh   rl   r   r   ru   ry   Ú__static_attributes__rp   r!   r   r   r      s¹   † ñð ;?Ø+/ð	HàðHð 8ðHð )õ	Hö&7ð$ ó&ó ð&ð ó8ó ð8ô1ô!ôDô
ôFô
ôEö
2ð óó ðô#ð õó ór!   r   c                  ó’   ^ • \ rS rSrSr    S         SU 4S jjjr  SS jrSS jrSSS jjrSS jr	SS jr
SS	 jrS
rU =r$ )ÚMMKernelInputséÒ   zz
Specialized KernelInputs for matrix multiplication operations.
Provides additional methods to access M, N, K dimensions.
c                óZ  >• [         TU ]  XU5        [        U R                  5      S:¼  d   S5       eXEpvUS:  a  U[        U5      -  nUS:  a  U[        U5      -  nSUs=::  a  [        U5      :  d  O   SU 35       eSUs=::  a  [        U5      :  d  O   SU 35       eX@l        XPl        g)z“
Initialize with a tuple of input nodes.

By default, we assume the last 2 input nodes are mat1 and mat2, but
the caller can adjust when necessary
é   zExpected at least 2 input nodesr   zInvalid mat1_idx: zInvalid mat2_idx: N)Úsuperr   r   r   Ú	_mat1_idxÚ	_mat2_idx)	r   r   r   r   Úmat1_idxÚmat2_idxÚm1_idxÚm2_idxÚ	__class__s	           €r   r   ÚMMKernelInputs.__init__Ø   s¼   ø€ ô 	‰Ñ˜¨yÔ9ô �4×$Ñ$Ó%¨Ó*ÐMÐ,MÓMÐ*ð "�Ø�a‹<Ø”c˜+Ó&Ñ&ˆFØ�a‹<Ø”c˜+Ó&Ñ&ˆFà�FÕ-œS Ó-Õ-ÐNÐ1CÀHÀ:Ð/NÓNÐ-Ø�FÕ-œS Ó-Õ-ÐNÐ1CÀHÀ:Ð/NÓNÐ-à!ŒØ!�r!   c                óh  • U R                  5       U R                     nU R                  5       U R                     nUR                  5       S   nUR                  5       S   nUR                  5       S   nUR                  5       S   n[        R
                  R                  R                  XF5        X5U4$ )aq  
Get the symbolic M, N, K dimensions for matrix multiplication.
Handles both 2D (MM) and 3D (BMM) tensors.

M is extracted from the second-to-last dimension of the first operand (mat1).
N is extracted from the last dimension of the second operand (mat2).
K is extracted from the last dimension of the first operand (mat1).

Returns:
    A tuple of (M, N, K) dimensions
éþÿÿÿéÿÿÿÿ)r&   r”   r•   r?   r   rN   rO   Úcheck_equals)r   Úmat1Úmat2ÚmÚkÚnÚk0s          r   Úmnk_symbolicÚMMKernelInputs.mnk_symbolicù   s�   € ð �z‰z‹|˜DŸN™NÑ+ˆØ�z‰z‹|˜DŸN™NÑ+ˆà�M‰M‹O˜BÑˆØ�M‰M‹O˜BÑˆØ�M‰M‹O˜BÑˆð �]‰]‹_˜RÑ ˆÜ	�‰×Ñ×%Ñ% aÔ,Ø�aˆyÐr!   c                óv   • U R                   b  U R                   $ U R                  5       S   R                  5       $ )ro   r   )r   Úmat1mat2rf   r)   s    r   r   ÚMMKernelInputs.out_dtype  s2   € ð �?‰?Ñ&Ø—?‘?Ð"Ø�}‰}‹˜qÑ!×+Ñ+Ó-Ð-r!   c                ó²  • U R                  5       u  p#U R                  5       nUR                  5       Gt pVnUR                  5       Gt p‰n
[        XX5       VVs/ s H.  u  p¼[        R
                  R                  R                  X¼5      PM0     snnn/ UQUPU
PnU(       a  [        U R                  5       XM5      $ [        U R                  5       XM5      $ s  snnf )zÐ
Handle output layout generation for matrix multiplication.

Args:
    out_dtype: Optional output dtype. If not provided, infer from inputs
    flexible: If True, return FlexibleLayout, otherwise FixedLayout
)r©   r   r?   Úzipr   rN   rO   Úcheck_equals_and_simplifyr   r2   r   )r   rx   r    r¡   r   Úb1r¢   Úk1Úb2Úk2r¤   ÚaÚbÚsizes                 r   ry   ÚMMKernelInputs.output_layout  sª   € ð —]‘]“_‰
ˆØ—N‘NÓ$ˆ	à—]‘]“_‰
ˆ�Ø—]‘]“_‰
ˆ�ÜJMÈbÌ+ÔVÊ+Á$À!ŒQ�W‰W×Ñ×7Ñ7¸Ö=É+ÒVˆØ�ˆz�Aˆz�qˆzˆÞÜ! $§+¡+£-°ÓAÐAä˜tŸ{™{›}¨iÓ>Ð>ùó Ws   Á5Cc                óZ   • U R                  5       nXR                     XR                     4$ )zJ
Get the mat1 and mat2 nodes.

Returns:
    A tuple of (mat1, mat2) nodes
)r&   r”   r•   )r   r&   s     r   r©   ÚMMKernelInputs.mat1mat22  s(   € ð —
‘
“ˆØ—^‘^Ñ$ e¯N©NÑ&;Ð;Ð;r!   c                ó®   • U R                  5       nXR                     nXR                     nUS   nUS   nUS   nUS   nXW:X  d   SU SU 35       eXFU4$ )zð
Get the hinted M, N, K dimensions for matrix multiplication.
Handles both 2D (MM) and 3D (BMM) tensors.

Uses shapes_hinted from the base class to get integer hints for dimensions.

Returns:
    A tuple of (M, N, K) dimensions as integers
r�   rž   zK dimensions don't match: r#   )rV   r”   r•   )r   Úhinted_shapesÚ
mat1_shapeÚ
mat2_shaper¢   r£   r¤   Úk_checks           r   Ú
mnk_hintedÚMMKernelInputs.mnk_hinted<  sw   € ð ×*Ñ*Ó,ˆØ"§>¡>Ñ2ˆ
Ø"§>¡>Ñ2ˆ
à�r‰NˆØ�r‰NˆØ�r‰Nˆð ˜R‘.ˆØ‹|ÐJÐ9¸!¸¸DÀÀ	ÐJÓJˆ|à�aˆyÐr!   c                óh   • U R                  5       nXR                     n[        U5      S:¼  a  US   $ g)z”
Get the hinted batch size for batched matrix multiplication.
Returns 1 for non-batched (2D) operations.

Returns:
    The batch size as an integer
é   éýÿÿÿr   )rV   r”   r   )r   r¹   rº   s      r   Úbatch_hintedÚMMKernelInputs.batch_hintedT  s7   € ð ×*Ñ*Ó,ˆØ"§>¡>Ñ2ˆ
äˆz‹?˜aÓØ˜b‘>Ð!Ør!   )r”   r•   )NNr�   rž   )
r   r{   r   r|   r   r}   r–   r€   r—   r€   )r~   z2tuple[sympy.Integer, sympy.Integer, sympy.Integer]r‚   r„   r…   )r~   ztuple[Any, Any])r~   ztuple[int, int, int]r   )r‡   rˆ   r‰   rŠ   r‹   r   r¦   r   ry   r©   r½   rÂ   r�   Ú__classcell__)rš   s   @r   r�   r�   Ò   s|   ø† ñð ;?Ø+/ØØð"àð"ð 8ð"ð )ð	"ð
 ð"ð ÷"ð "ðBà	;ôô4	.ö?ô(<ô÷0ò r!   r�   )Ú
__future__r   Úabcr   r   Útypingr   r   r   r	   r6   Útorch._inductor.configÚtorch._inductorr
   Útorch._inductor.virtualizedr   r   r   r   Úcollections.abcr   Úsympyr   r�   rp   r!   r   Ú<module>rÍ      sK   ðÝ "ç #ß 6Ó 6ã Û Ý Ý )ç 3Ñ 3ö Ý(ãô{�3ô {ô|O�\õ Or!   