Ë
    g^(h^   ã                   óÄ  — d dl Z d dlZd dlmZ d dlZd dlZddlmZ ddlm	Z	m
Z
 ddlmZ  G d„ de«      Z G d	„ d
e«      Ze j                  defd„«       Zdej"                  defd„Zdej"                  defd„Zdej"                  defd„Z G d„ de«      Z G d„ de«      Z G d„ de«      ZdgdggZdgdggdgdggdgdgggZg d¢g d¢g d¢gZdej"                  defd„Zy) é    N)ÚIntEnumé   )Úir)Úget_dtype_sizeÚsympy_product)ÚVc                   ó   — e Zd ZdZdZdZy)Ú	NCCL_COLLr   r   é   N)Ú__name__Ú
__module__Ú__qualname__Ú
ALL_REDUCEÚ
ALL_GATHERÚREDUCE_SCATTER© ó    ú[/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/_inductor/comm_analysis.pyr
   r
      s   „ Ø€JØ€JØ�Nr   r
   c                   ó   — e Zd ZdZdZdZy)ÚNVIDIA_GPU_TYPEr   r   r   N)r   r   r   ÚVOLTAÚAMPEREÚHOPPERr   r   r   r   r      s   „ Ø€EØ€FØ�Fr   r   Úreturnc                  ó8  — t         j                  j                  j                  t         j                  j                  j                  «      xs d} d| v rt
        j                  S d| v rt
        j                  S d| v rt
        j                  S t
        j                  S )NÚ ÚV100ÚA100ÚH100)	ÚtorchÚutilsÚcollect_envÚget_gpu_infoÚrunr   r   r   r   )Úgpu_infos    r   Úget_gpu_typer&      s|   € ä�{‰{×&Ñ&×3Ñ3´E·K±K×4KÑ4K×4OÑ4OÓPÒVÐTV€HØ�ÑÜ×$Ñ$Ð$Ø	�8Ñ	Ü×%Ñ%Ð%Ø	�8Ñ	Ü×%Ñ%Ð%ô ×%Ñ%Ð%r   Únodec                 ó  — t        | t        j                  «      st        d| › �«      ‚| j                  }|€J ‚d|v rt
        j                  S d|v rt
        j                  S d|v rt
        j                  S t        d|› �«      ‚)Nz!node is not a collective kernel: Ú
all_reduceÚ
all_gatherÚreduce_scatterzUnsupported collective kernel: )	Ú
isinstancer   Ú_CollectiveKernelÚ
ValueErrorÚpython_kernel_namer
   r   r   r   )r'   Úkernel_names     r   Úget_collective_typer1   (   s‹   € Ü�dœB×0Ñ0Ô1ÜÐ<¸T¸FÐCÓDÐDà×)Ñ)€KØÐ"Ð"Ð"Ø�{Ñ"Ü×#Ñ#Ð#Ø	˜Ñ	$Ü×#Ñ#Ð#Ø	˜[Ñ	(Ü×'Ñ'Ð'äÐ:¸;¸-ÐHÓIÐIr   c                 óV  — d}| j                   D ]—  }t        |j                  j                  «      }t	        |t
        j                  «      rt        |«      }n+t        j                  j                  j                  |d¬«      }||t        |j                  j                  «      z  z  }Œ™ |S )Nr   )Úfallback)Úinputsr   ÚlayoutÚsizer,   ÚsympyÚIntegerÚintr   ÚgraphÚsizevarsÚ	size_hintr   Údtype)r'   Úsz_bytesÚinpÚnumels       r   Úget_collective_input_size_bytesrA   8   s‡   € Ø€HØ�{‰{ò =ˆÜ˜cŸj™jŸo™oÓ.ˆÜ�eœUŸ]™]Ô+ä˜“J‰Eä—G‘G×$Ñ$×.Ñ.¨u¸qÐ.ÓAˆEØ�EœN¨3¯:©:×+;Ñ+;Ó<Ñ<Ñ<‰ð=ð €Or   c                 óŒ   — t        | «      t        j                  k(  rddlm}  || j
                  d   «      S t        d| › �«      ‚)Nr   )Ú_get_group_size_by_nameéÿÿÿÿzUnsupported collective type: )Útyper   r-   Ú"torch.distributed.distributed_c10drC   Úconstant_argsÚ	TypeError)r'   rC   s     r   Úget_collective_group_sizerI   E   s@   € ÜˆDƒz”R×)Ñ)Ò)ÝNá& t×'9Ñ'9¸"Ñ'=Ó>Ð>äÐ7¸°vÐ>Ó?Ð?r   c                   ó   — e Zd ZdZdZdZy)ÚNCCL_HWr   r   r   N)r   r   r   ÚNVLINKÚPCIÚNETr   r   r   rK   rK   S   s   „ Ø€FØ
€CØ
�Cr   rK   c                   ó   — e Zd ZdZdZy)Ú	NCCL_ALGOr   r   N)r   r   r   ÚTREEÚRINGr   r   r   rP   rP   Y   s   „ Ø€DØ�Dr   rP   c                   ó   — e Zd ZdZy)Ú
NCCL_PROTOr   N)r   r   r   ÚLLr   r   r   rT   rT   ^   s	   „ ð 
�Br   rT   g333333@gffffff@g333333ã?ç      ð?g      @gš™™™™™@)ç     €C@rW   gffffff4@)gÍÌÌÌÌìU@g     €6@g      3@c                 ój  — t        | «      }|dz  dz  dz  }d}t        | «      }t        j                  ||z  «      }|}|dk  ryt        j
                  }t        j                  }t        | «      }	t        j                  j                  j                  }
t        j                  j                  j                  }t        «       }|dk  r|dz
  nd}|dk(  r|nd}t        |   |   }|dk(  r|
n|}d}||z  }t!        |||dkD  s|	t"        j$                  k(  rdndz  «      }|	t"        j$                  k(  r	d|dz
  z  }n'|	t"        j&                  t"        j(                  fv r|dz
  }d|z  z  }||z  }|d	z  }t*        j,                  }|	t"        j$                  k(  r|dkD  rd|z  }n*d}n'|	t"        j&                  t"        j(                  fv r|dz
  }t.        |   |   }t0        |   |   |   }t0        t*        j2                     |   |   }d
}|dkD  rd}t5        ||«      }||z
  |z  ||z  z   z  }|dz  }||z  }||z   S )a9  
    Returns estimated NCCL collective runtime in nanoseconds (ns).

    The following heuristics are copied from https://github.com/NVIDIA/nccl/blob/master/src/graph/tuning.cc.
    We aim to estimate the runtime as accurately as possible.

    Assumptions:
    - only ring algorithm (NCCL_ALGO_RING) is used
    - only Low-Latency protocol (NCCL_PROTO_LL) is used, i.e. Simple or LL128 is not used
    - 8 gpus per node  # TODO: Need to find a way to get accurate "gpus per node" and "# nodes" info.
    - collective is one of: allreduce, reducescatter, allgather
    i   é   r   r   r   g      Ð?gUUUUUUÕ?rV   g    eÍÍAg        g     @�@)rA   rI   ÚmathÚceilrP   rR   rT   rU   r1   r    Ú	_inductorÚconfigÚintra_node_bwÚinter_node_bwr&   ÚllMaxBwsÚminr
   r   r   r   rK   rL   ÚbaseLatÚhwLatrN   Úmax)r'   Útensor_storage_size_bytesÚtensor_storage_size_GBÚnum_gpus_per_nodeÚ
group_sizeÚnNodesÚnRanksÚ	nccl_algoÚ
nccl_protoÚcollÚbwIntraÚbwInterÚcompCapIndexÚindex2Úindex1ÚllMaxBwÚbwÚ	nChannelsÚbusBwÚnstepsÚratioÚ	bandwidthÚbandwidth_GB_per_nsÚintraHwÚnInterStepsÚlatencyÚintraLatÚinterLatÚnetOverheadÚ
latency_nsÚtransport_nss                                  r   Ú estimate_nccl_collective_runtimerƒ   ¡   sr  € ô !@ÀÓ EÐà6¸Ñ=ÀÑDÀtÑKÐð ÐÜ*¨4Ó0€JÜ�Y‰Y�zÐ$5Ñ5Ó6€FØ€Fà�‚{Øô —‘€IÜ—‘€JÜ˜tÓ$€Dô
 �o‰o×$Ñ$×2Ñ2€GÜ�o‰o×$Ñ$×2Ñ2€Gä“>€LØ! Qš;ˆV�aŠZ¨A€Fà# qš[‰\¨a€FÜ�vÑ˜vÑ&€Gð ˜a’K‰ W€BØ€IØ˜‰N€Eô ØØØ !š t¬y×/CÑ/CÒ'C‰9È)ñ	Uó€Eð Œy×#Ñ#Ò#Ø�f˜q‘jÑ!‰Ø	”)×*Ñ*¬I×,@Ñ,@ÐAÑ	AØ˜!‘ˆð �6‰\˜VÑ#€EØ˜‘€Ià# c™/Ðô �n‰n€GàŒy×#Ñ#Ò#Ø�AŠ:Ø˜f™*‰Kà‰KØ	”)×*Ñ*¬I×,@Ñ,@ÐAÑ	AØ˜q‘jˆô �iÑ  Ñ,€GÜ�W‰~˜iÑ(¨Ñ4€HÜ”W—[‘[Ñ! )Ñ,¨ZÑ8€Hð €KØ�‚zØˆÜ�8˜[Ó)€HØ�˜Ñ$¨Ñ0°;ÀÑ3IÑIÑI€Gà˜3‘€Jð *Ð,?Ñ?€LØ˜*Ñ$Ð$r   )Ú	functoolsrZ   Úenumr   r7   r    r   r   r!   r   r   Úvirtualizedr   r
   r   Ú	lru_cacher&   ÚIRNoder1   r9   rA   rI   rK   rP   rT   rb   rc   r`   Úfloatrƒ   r   r   r   ú<module>rŠ      sJ  ðÛ Û Ý ã ã å ß 0Ý ô�ô ô�gô ð ×Ñð
&�oò 
&ó ð
&ðJ˜bŸi™ið J¨Ió Jð 
¨"¯)©)ð 
¸ó 
ð@ B§I¡Ið @°#ó @ôˆgô ô�ô ô
�ô ð 	ðð
 	ðð	€ð  
ˆØ	ˆðð 
ˆØ	ˆðð 
ˆØ	ˆðð	€ò,òòð€ð,b%¨2¯9©9ð b%¸ô b%r   