
    ^jE                    `	   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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Zd dlmZ d dlmZ d dlmZ d dlmZ d dlmZ d dlmZmZmZ d dlmZ d d	l m!Z!m"Z" d d
l#m$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. d dl/m0Z0 d dl1m2Z2 d dl3m4Z4 d dl5m6Z6 d dl7m8Z8m9Z9m:Z:m;Z;m<Z<m=Z=m>Z>m?Z?m@Z@mAZAmBZBmCZCmDZD d dlEmFZFmGZGmHZH dZIdZJdZKdZLdZMdZNe4r+d dlOZOdD ]"  \  ZPZQ eOj                  eP      sd eS       eQ<   $ deTdeUdeUde%j                  deTdeTde*j                  fdZX ej                  eZ      Z[e[j                  ej                         g dZ^ddgZ_e?xs e@xs eCZ` G d  d!e"      Zai d" ead#d$      d% ead&d'      d( ead)d*      d+ ead,d-      d. ead/d0      d1 ead2d3      d4 ead5d6      d7 ead8d9      d: ead;d<      d= ead>d?      d@ eadAdB      dC eadDdE      dF eadGdH      dI eadJdK      dL eadMdN      dO eadPdQ      dR eadSdT      dU eadVdW      iZbe G dX dY             ZcdZ Zdd[ Zed\ Zfd] Zgd^ Zhd_ Zid` ZjdekfdaZlddbdcZmddeUfdeZndfe!dgeUdeUfdhZodi Zpdj Zqdk Zrdl Zsdm Ztdn Zudo Zvdp Zwdq Zxdr ZyddsZzdt Z{du Z| e;dv       Z}eke~dw<   dx ZdyeeTdzf   fd{Zdd|Zd} Zde%j                  d~eUdeUdekfdZe<ddd ed      dddfd       ZeBrdZn eU ej                  dd            ZdddZeArded<   ddekfdZdeUfdZed        ZddeUdeUdeUfdZdeUdeTfdZdaej                  dz  e~d<   ddeTdz  ddfdZddZ	 dde%j
                  j$                  de%j
                  j&                  deUdz  de!fdZeAr+ ed ede%j.                  j1                                     ZndZ G d deD      Z G d de      ZdeeTee!   f   dede!fdZej>                  dekfd       Zd ZdeefdZ G d deD      Z G d de,jH                        Z G d de,jH                        Ze	 dd       Z G d de%jP                  jR                  j                        Z G d de      Z G d deD      Z G d de      Zy)    N)Callable)contextmanager)	dataclass)	timedelta)Enum)partialreducewraps)StringIO)Any
NamedTuple)patch)
DeviceType)_SymmetricMemory)	trace_log)_TORCHCOMM_AVAILABLE)common_utils)FILE_SCHEMAfind_free_portIS_SANDCASTLELazyValretry_on_connect_failuresskip_but_pass_in_sandcastleskip_but_pass_in_sandcastle_if	TEST_CUDATEST_HPUTEST_WITH_ROCMTEST_WITH_TSANTEST_XPUTestCase)_install_threaded_pg_uninstall_threaded_pgProcessLocalGroupF))glooTORCHCOMM_HAS_GLOO)xcclTORCHCOMM_HAS_XCCL)ncclTORCHCOMM_HAS_NCCL)rcclxTORCHCOMM_HAS_RCCLX)ncclxTORCHCOMM_HAS_NCCLXTbackendrank
world_sizedevice
store_path
group_namereturnc                 "   ddl m}m} t        j                  ||      }t        |      t        j                  d<   t        |      t        j                  d<   t        j                  | || dt        j                  | d|      ddi	      }	t        j                  t        j                  | d
|      |	j                         |	j                               }
|
j                  |	j                         t        j                  j                   j"                   ||	             |
j%                  t        j                  j                   j"                         |
j'                  |        ||
| ||| t)        |      D ci c]  }|| c}       |
S c c}w )ur  Build a bare ProcessGroup whose backend is a torchcomms _BackendWrapper.

    Creates a torchcomms comm of the requested ``backend`` (e.g. ``nccl``,
    ``ncclx``) and wraps it in ``_BackendWrapper``, then registers it as the
    NCCL backend on a fresh ``ProcessGroup`` and publishes the group into
    ``c10d._world`` so name-based lookups (e.g. ``symm_mem.rendezvous(...,
    group=group_name)``) resolve to it.

    Callers are responsible for choosing a ``group_name`` (and matching
    ``store_path``) that is unique within the process — duplicates trip
    ``_register_pg_in_world`` and the underlying FileStore.
    r   )_BackendWrapper_register_pg_in_worldTORCHCOMM_RANKTORCHCOMM_SIZE_commz_tc/persistent_storetrue)namestorehintsz_pg/)backend_namer>   r3   backend_configrank_mapping)"torch.distributed.distributed_c10dr6   r7   c10d	FileStorestrosenviron
torchcommsnew_commPrefixStoreProcessGroupget_rankget_size_register_backend
get_deviceBackendTypeNCCL_set_default_backend_set_group_namerange)r.   r/   r0   r1   r2   r3   r6   r7   
file_storetc_commpgrs               u/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/testing/_internal/common_distributed.pysetup_torchcomms_pgr[   R   sb   (
 
J7J#&t9BJJ #&z?BJJ !!|5!*T2J?!6*G 
		J<t,j9
B
 %%**  D--99>>?z"
$)*$56qad6 I 7s   8
Fr(   r&   hcclcudaxpuc                   "    e Zd ZU eed<   eed<   y)TestSkip	exit_codemessageN)__name__
__module____qualname__int__annotations__rF        rZ   ra   ra      s    NLrj   ra   backend_unavailableH   z5Skipped because distributed backend is not available.small_worldsizeI   z Skipped due to small world size.odd_worldsizeW   zSkipped due to odd world size.no_cudaJ   zCUDA is not available.zmulti-gpu-1K   zNeed at least 1 CUDA devicezmulti-gpu-2M   zNeed at least 2 CUDA deviceszmulti-gpu-3P   zNeed at least 3 CUDA deviceszmulti-gpu-4Q   zNeed at least 4 CUDA deviceszmulti-gpu-5R   zNeed at least 5 CUDA deviceszmulti-gpu-6S   zNeed at least 6 CUDA deviceszmulti-gpu-7T   zNeed at least 7 CUDA deviceszmulti-gpu-8U   zNeed at least 8 CUDA devicesr(   L   z#c10d not compiled with NCCL support
skipIfRocmN   zTest skipped for ROCmno_peer_accessO   z'Test skipped because no GPU peer accessgenericV   zHTest skipped at subprocess level, look at subprocess log for skip reasonimporterrorX   z"Test skipped due to missing importno_acceleratorY   zaccelerator is not available.c                       e Zd Zi Zh ded<    e       ed<   h ded<   h ded<   i Zh ded<   h ded	<   h ded
<   h ded<    e       ed<   erdhed<   erdhed<   yy)DistTestCases>   mpiuccr(   r&   allgather_coalescedr	   >   r   r(   r&   zsendrecv anysourcezcpu barrier>   r   r$   r(   gpur^   ddpsubgrouppluginr]   hpur&   r_   N)rd   re   rf   skip_collectivesetbackend_featurer   r   ri   rj   rZ   r   r      s     O-KO)* #OH,CO()%<OM" O4OE5OF4OE"9OJ #OH"("( rj   r   c                     | t         v S N)DDP_RANK_DEVICESr1   s    rZ   requires_ddp_rankr      s    %%%rj   c                 .     t                fd       }|S )zSkips if the world size exceeds the number of GPUs, ensuring that if the
    test is run, each rank has its own GPU via ``torch.cuda.device(rank)``.c                     t         s2t        s,t        s&t        j                  t
        d   j                         t        t        j                  d         }t         rJt        j                  j                         |k  r)t        j                  t
        d|    j                         t        rJt        j                  j                         |k  r)t        j                  t
        d|    j                         t        rJt        j                  j                         |k  r)t        j                  t
        d|    j                          | i |S )Nrq   
WORLD_SIZE
multi-gpu-)r   r   r   sysexit
TEST_SKIPSrb   rg   rG   rH   torchr^   device_countr   r_   )argskwargsr0   funcs      rZ   wrapperzskip_if_no_gpu.<locals>.wrapper   s    XHHZ	*445L12
002Z?HHZ*ZL 9:DDE		..0:=HHZ*ZL 9:DDE		..0:=HHZ*ZL 9:DDET$V$$rj   r
   r   r   s   ` rZ   skip_if_no_gpur      s"     4[% % Nrj   c                 .     t                fd       }|S )Nc                      t         j                  d   dk7  rEt        t         j                  d         dk  r&t        j                  t
        d   j                          | i |S )NBACKENDr   r      rm   rG   rH   rg   r   r   r   rb   r   r   r   s     rZ   r   z(skip_if_small_worldsize.<locals>.wrapper   sR    JJy!U*BJJ|4L0MPQ0QHHZ 12<<=T$V$$rj   r   r   s   ` rZ   skip_if_small_worldsizer           
4[% % Nrj   c                 .     t                fd       }|S )Nc                      t         j                  d   dk7  rHt        t         j                  d         dz  dk(  r&t        j                  t
        d   j                          | i |S )Nr   r   r         ro   r   r   s     rZ   r   z&skip_if_odd_worldsize.<locals>.wrapper   sW    JJy!U*BJJ|4L0MPQ0QUV0VHHZ0::;T$V$$rj   r   r   s   ` rZ   skip_if_odd_worldsizer      r   rj   c                       fd}|S )Nc                 4     t                fd       }|S )Nc                      dk(  rKt         j                  j                         k  r*t        j                  t
        d    j                         y  | i |S Nr(   r   )r   r^   r   r   r   r   rb   )r   r   r.   r   ns     rZ   r   zCrequire_n_gpus_for_nccl_backend.<locals>.decorator.<locals>.wrapper  sM    & UZZ%<%<%>%Bj$45??@T,V,,rj   r   )r   r   r.   r   s   ` rZ   	decoratorz2require_n_gpus_for_nccl_backend.<locals>.decorator  s     	t	- 
	- rj   ri   )r   r.   r   s   `` rZ   require_n_gpus_for_nccl_backendr     s     rj   c                      d } | S )Nc                 .     t                fd       }|S )Nc                      	 ddl m}m}  | i |S # t        $ r) t	        j
                  t        d   j                         Y y w xY w)Nr   )AutoModelForMaskedLM
BertConfigr   )transformersr   r   ImportErrorr   r   r   rb   )r   r   r   r   r   s       rZ   r   z?import_transformers_or_skip.<locals>.decorator.<locals>.wrapper  sA    >IT,V,, >M2<<=>s    /AAr   r   s   ` rZ   r   z.import_transformers_or_skip.<locals>.decorator  s     	t	> 
	> rj   ri   )r   s    rZ   import_transformers_or_skipr     s    
 rj   c                     t         r"t        j                  j                         | k\  ryt        r"t        j
                  j                         | k\  ryt        r"t        j                  j                         | k\  ryyNTF)r   r   r^   r   r   r   r   r_   )xs    rZ   at_least_x_gpur   $  sS    UZZ,,.!3EII**,1EII**,1rj   c                 V    t        | d   dd       }t        |       dk(  s|y ||       y)Nr   _handle_test_skipFT)getattrlen)r   msgr   s      rZ   _maybe_handle_skip_if_lt_x_gpur   .  s5    Q)<dC
4yA~*2crj   )	allow_cpuc                      fd}|S )zSkip if fewer than x accelerators available.

    Args:
        x: Minimum number of accelerators required.
        allow_cpu: If True, run the test on CPU-only machines (no accelerators).
    c                 4     t                fd       }|S )Nc                  <   t         j                  j                         r)t         j                  j                         k\  r | i |S t        r)t         j
                  j                         k\  r | i |S t        r)t         j                  j                         k\  r | i |S r2t         j                  j                         st        st        s | i |S t        d    }t        | |j                        s t        j                  |j                         y y )Nr   )r   r^   is_availabler   r   r   r   r_   r   r   rc   r   r   rb   )r   r   	test_skipr   r   r   s      rZ   r   z4skip_if_lt_x_gpu.<locals>.decorator.<locals>.wrapper?  s    zz&&(UZZ-D-D-F!-KT,V,,EII2249T,V,,EII2249T,V,,%**"9"9";x8T,V,,"Zs#34I1$	8I8IJ,,- Krj   r   )r   r   r   r   s   ` rZ   r   z#skip_if_lt_x_gpu.<locals>.decorator>  s     	t	. 
	. rj   ri   )r   r   r   s   `` rZ   skip_if_lt_x_gpur   6  s    " rj   r   c                       fd}|S )aR  
    Decorator to request a specific world size for a test. The test harness can
    read this attribute to set the number of ranks to spawn. If there are fewer
    than `n` CUDA devices available, the test should be skipped by the harness.

    Usage:
        @require_world_size(3)
        def test_something(self):
            ...
    c                     | _         t        j                  j                         } t	        j
                  |k\  d d|       |       S )Nz	requires z GPUs, found )_required_world_sizer   r^   r   unittest
skipUnless)r   	availabler   s     rZ   r   z&requires_world_size.<locals>.decorator^  sR    $%!JJ++-	
x""Nis-	{C

 	rj   ri   )r   r   s   ` rZ   requires_world_sizer   R  s     rj   objdefaultc                     	 t        | d      r%t        | j                        r| j                         n| j                  }t	        | |      }|j
                  }t        |      S # t        $ r |cY S w xY w)z
    Returns the requested world size for the currently running unittest method on `obj`
    if annotated via `@require_world_size(n)`, else returns `default`.
    _current_test_name)hasattrcallabler   _testMethodNamer   r   rg   	Exception)r   r   	test_namefnvalues        rZ   get_required_world_sizer   h  st    
 s01hs?U?U6V ""$$$ 	
 S)$''5z s   AA" "A0/A0c                       fd}|S )Nc                 4     t                fd       }|S )Nc                  2   dk7  r | i |S t         j                  j                         r)t         j                  j                         k\  r | i |S t        d    }t        | |j                        s t        j                  |j                         y y r   )
r   r^   r   r   r   r   rc   r   r   rb   )r   r   r   r.   r   r   s      rZ   r   z9nccl_skip_if_lt_x_gpu.<locals>.decorator.<locals>.wrapper~  s    & T,V,,zz&&(UZZ-D-D-F!-KT,V,,"Zs#34I1$	8I8IJ,,- Krj   r   )r   r   r.   r   s   ` rZ   r   z(nccl_skip_if_lt_x_gpu.<locals>.decorator}  s     	t	. 
	. rj   ri   )r.   r   r   s   `` rZ   nccl_skip_if_lt_x_gpur   |  s     rj   c                    | j                         }d|vrt        d      d|vrt        d      d|vrt        d      |d   }|j                  d      dk(  r|n|j                  d      d	   }||vrt        d
| d|       y )N	iterationz(Expected 'iteration' in ddp_logging_data	has_errorz(Expected 'has_error' in ddp_logging_dataerrorz$Expected 'error' in ddp_logging_dataz
Exception raised from r   zDid not find expected z in ddp logging data error: )_get_ddp_logging_dataAssertionErrorfindsplit)	model_DDP
err_substrddp_logging_datalogging_erractuals        rZ   verify_ddp_error_loggedr     s     668**GHH**GHH&&CDD"7+K ??56"< 	89!< 
 [ $VH,HV
 	
 !rj   c                 .     t                fd       }|S )aJ  
    Convenience decorator to set/unset TORCH_NCCL_BLOCKING_WAIT flag. Note that use of
    this decorator will override the setting of TORCH_NCCL_ASYNC_ERROR_HANDLING for
    the particular test. After the test, both TORCH_NCCL_BLOCKING_WAIT and
    TORCH_NCCL_ASYNC_ERROR_HANDLING will be restored to their original values.
    c                     	 t         j                  d   }t         j                  d= 	 t         j                  d   }dt         j                  d<   	  | i |}|||t         j                  d<   ||t         j                  d<   S S # t        $ r d }Y jw xY w# t        $ r d }Y gw xY w# dt         j                  d<   w xY w# ||t         j                  d<   ||t         j                  d<   w w xY w)NTORCH_NCCL_ASYNC_ERROR_HANDLINGTORCH_NCCL_BLOCKING_WAIT1)rG   rH   KeyError)r   r    cached_nccl_async_error_handlingcached_nccl_blocking_waitretr   s        rZ   r   z(with_nccl_blocking_wait.<locals>.wrapper  s   	4;=::1<, 

<=	946JJ*5% 69BJJ12	S''C 0;4 

<= )49R

56 51  	4/3,	4  	-(,%	- 69BJJ12 0;4 

<= )49R

56 5s@   $B B 	B> BBB# B& "B##B& &B;>-C+r   r   s   ` rZ   with_nccl_blocking_waitr    s%     4[ S  SD Nrj   c                       fd}|S )zK
    Runs a test for each distributed debug level specified in levels.
    c                 2     t                fd       }|S )Nc                     t         j                  j                  dd       }D ][  }|t         j                  d<   t        j                           | i |}t        j
                          |I|t         j                  d<   ] S )NTORCH_DISTRIBUTED_DEBUG)rG   rH   getrD   set_debug_level_from_envbarrier)r   r   	old_levellevelr  r   levelss        rZ   r   z:with_dist_debug_levels.<locals>.decorator.<locals>.wrapper  sx    

'@$GI F8=

45--/D+F+(<EBJJ89F Jrj   r   )r   r   r  s   ` rZ   r   z)with_dist_debug_levels.<locals>.decorator  s     	t	 
	 rj   ri   )r  r   s   ` rZ   with_dist_debug_levelsr    s    
$ rj   c                  @    t        t        j                          d      S )Nz+c10d was not compiled with the Gloo backend)r   rD   is_gloo_availableri   rj   rZ   requires_gloor    !    )""$$5 rj   c           	         t         sd S t        j                         st        d      S t	        t
        j                  j                  j                         | k  d|  dt
        j                  j                  j                          d|       S )Nc                     | S r   ri   )fs    rZ   <lambda>z'requires_nccl_version.<locals>.<lambda>  s     rj   +c10d was not compiled with the NCCL backendz0Requires NCCL version greater than or equal to: z	, found: z
, reason: )	r   rD   is_nccl_availabler   r   r   r^   r(   version)r  r   s     rZ   requires_nccl_versionr    s    !!#*9
 	
 .JJOO##%/>wiyQVQ[Q[Q`Q`QhQhQjPkkuvyuz{
 	
rj   c                      t        dd      S )zK
    Require NCCL shrink support (NCCL available and version >= 2.27).
    )r      z Need NCCL 2.27+ for shrink_group)r  ri   rj   rZ   requires_nccl_shrinkr     s     !*LMMrj   c                  @    t        t        j                          d      S )Nr  )r   rD   r  ri   rj   rZ   requires_ncclr     r  rj   c                  @    t        t        j                          d      S )Nz*c10d was not compiled with the UCC backend)r   rD   is_ucc_availableri   rj   rZ   requires_uccr#    !    )!!##4 rj   c                  @    t        t        j                          d      S )Nz*c10d was not compiled with the MPI backend)r   rD   is_mpi_availableri   rj   rZ   requires_mpir'    r$  rj   c                 V    | t         } t        d | D              }t        | d|        S )a  
    Decorator to skip tests if no accelerator communication backend (NCCL, XCCL, HCCL) is available.

    Args:
        backends (Optional[List[str]]): Specific accelerator backends to check (e.g., ["nccl", "xccl", "hccl"]).
                                       If None, checks all supported accelerator backends (NCCL, XCCL, HCCL).

    Returns:
        callable: A decorator that skips the test if no specified accelerator backend is available.
    c              3      K   | ]=  }	 t        j                  t         j                  d  dj                  |d               ? yw)c                      t         S r   )r   ri   rj   rZ   r  z=requires_accelerator_dist_backend.<locals>.<genexpr>.<lambda>.  s    H rj   r\   c                       yNFri   ri   rj   rZ   r  z=requires_accelerator_dist_backend.<locals>.<genexpr>.<lambda>/  s    rj   N)rD   r  is_xccl_availabler	  ).0r.   s     rZ   	<genexpr>z4requires_accelerator_dist_backend.<locals>.<genexpr>*  sG       	&****$	
 #g}
%		(s   AAz5No accelerator communication backend available among )ACCELERATOR_DIST_BACKENDSanyr   )backendsbackend_availables     rZ   !requires_accelerator_dist_backendr4    sH     ,     *
?zJ rj   c                      t         j                  j                         xr$ t        j                  t
        j                  d      } t        |  d      S )Nr   z"multicast support is not available)r   r^   r   r   has_multicast_supportr   CUDAr   )r6  s    rZ   requires_multicast_supportr8  9  sI    

! 	G22:??AF  *!!, rj   c                      t         r@t        r9ddg} | D ]/  }|t        j                  j	                  d      j
                  v s/ y yyy)Ngfx942gfx950r   TF)r   r   r   r^   get_device_propertiesgcnArchName)	arch_listarchs     rZ   #evaluate_platform_supports_symm_memr@  D  sL    !8,I!  5::;;A>JJJ  rj   c                      t               S r   )r@  ri   rj   rZ   r  r  S  s
    /1 rj   PLATFORM_SUPPORTS_SYMM_MEMc                 d     t        j                  t        t        d   j                        |       S )z&Skips a test for ROCm multiprocess UTsr|   )r   skipIfr   r   rc   )r   s    rZ   skip_if_rocm_multiprocessrE  W  s%    L8??>:l+C+K+KLTRRrj   r?  .c                       fd}|S )z4Skips a test for given ROCm archs - multiprocess UTsc                     d }t         rDt        j                  j                  d      j                  j                  d      d   }|v rd } t        j                  |d u|      |       S )Nr   :z0skip_if_rocm_arch_multiprocess: test skipped on )r   r   r^   r<  r=  r   r   rD  )r   reasonpropr?  s      rZ   r   z1skip_if_rocm_arch_multiprocess.<locals>.decorator_  se    ::33A6BBHHMaPDt|KD6R:xvT16:4@@rj   ri   )r?  r   s   ` rZ   skip_if_rocm_arch_multiprocessrK  \  s    A rj   c                       fd}|S )z:Skips a test for ROCm based on ROCm ver - multiprocess UTsc                 :   d }t         rut        t        j                  j                        }|j                  dd      d   }t        d |j                  d      D              }||t              k  r	d| d d	} t        j                  |d u|      |       S )
N-r   maxsplitr   c              3   2   K   | ]  }t        |        y wr   )rg   )r.  r   s     rZ   r/  zLskip_if_rocm_ver_lessthan_multiprocess.<locals>.decorator.<locals>.<genexpr>s  s     &O!s1v&Os   .z-skip_if_rocm_ver_lessthan_multiprocess: ROCm z is available but z	 required)	r   rF   r   r  hipr   tupler   rD  )r   rI  rocm_versionrocm_version_tupler  s       rZ   r   z9skip_if_rocm_ver_lessthan_multiprocess.<locals>.decoratorn  s    u}}001L'--cA->qAL!&&O|7I7I#7N&O!O"*?%g6HI[H\\novnw  xA  B:xvT16:4@@rj   ri   )r  r   s   ` rZ   &skip_if_rocm_ver_lessthan_multiprocessrW  k  s    A rj   c                  <    t        t        j                  dk(  d      S )Nwin32z8This unit test case is not supported on Windows platform)r   r   platformri   rj   rZ   skip_if_win32r[    s    )B rj   majorminorc                     | j                   dk7  ryt        j                  j                  yt        j                  j                  |       ||fk\  S )z
    Returns True if the device's compute capability is (major, minor) or higher.
    Error out if the device is not a CUDA device.
    Returns False if device is a RoCM device.
    Returns True if device is a non-CUDA device.
    r^   TF)typer   r  rS  r^   get_device_capability)r1   r\  r]  s      rZ   sm_is_or_higher_thanra    sD     {{f}}$::++F3u~EErj   	localhostr      )minutesc                     t               }|rEt        |t        d      z        }t        j                  j
                  j                  | ||||      S t        j                  | |||||      S )zL
    Creates a TCP store. Retries if the chosen port is already in use.
    r   )milliseconds)wait_for_workers	use_libuv)r   rg   r   r   classes	dist_c10dTCPStorerD   )	addrr0   	is_mastertimeoutrg  	jit_classrh  porttimeout_milliseconds	            rZ   create_tcp_storerr    sr     D!'I1,E"EF}}&&//$
I/B
 	
 }}-
 	
rj   i  !DISTRIBUTED_TESTS_DEFAULT_TIMEOUT300i  )test_ddp_uneven_inputstest_DistributedDataParallel   test_join_kwargs	lazy_initc                     t         j                  dk(  s| !t        j                  j	                  d|      S t        j                  j	                  | |      S )NrY  z	127.0.0.1)hostnamery  	interfacery  )r   rZ  rD   ProcessGroupGloocreate_devicer|  s     rZ   r  r    s[    
||w)"3$$22 I 3 
 	
 $$229 3 
 	
rj   c                 Z    t         j                  | j                  d      d   t              S NrR  r   )TIMEOUT_OVERRIDEr	  r   TIMEOUT_DEFAULT)test_ids    rZ   get_timeoutr    s#    c 22 6HHrj   c               #   N  K   t               t               }} t        j                  t        j                  }}	 | |ct        _        t        _        t        j                  t        j                  f ||ct        _        t        _        y # ||ct        _        t        _        w xY wwr   )r   r   stdoutstderr)new_outnew_errold_outold_errs       rZ   captured_outputr    sl     z8:WGzz3::WG2!('
CJjj#**$$!('
CJ'
CJs   5B%9B	 1B%	B""B%
num_inputsc                    ddt         dt         dt         dt         fd}dt         fd}t        |d      t        |d	      t        |d
      t        |d      t        |d	      t        |d
      fD cg c]N  }t        |      D cg c]  } ||| z  |z   ||z         c}t        |      D cg c]  } ||||z         c}fP c}}S c c}w c c}w c c}}w )z
    Generate a number of basic test cases for sparse reduction.
    These cover tensors with a varying number of sparse dimensions and a varying
    number of dense dimensions. The only reduction operation we support is sum.
    r   r/   r0   sparse_dims
dense_dimsc           	         t        j                  t        j                  | dz         d| dz   f      }|gt        |      D cg c]  }d c}z   }t        |dz
        D ]A  }t        j                  |t        j
                  d| dz         f      }|j                  |       C t        j                  | dz   gt        |      D cg c]  }d c}z         }t        j                  |||      S c c}w c c}w )Nr   r   )	r   reshapearangerU   catzerosappendonessparse_coo_tensor)r/   r0   r  r  indices_shapevaluess           rZ   generatez,simple_sparse_reduce_tests.<locals>.generate  s     --TAX 6D1HF5+<=a=={Q' 	%Aii%++a*B CDGLL$	% TAXJU:5F)G!)GGH&&w>>  > *Hs   	C+	C0
c           
      |    t        t        j                  t        |      D cg c]  } | ||       c}      S c c}w r   )r	   operatoraddrU   )r   r0   r/   s      rZ   compute_sumz/simple_sparse_reduce_tests.<locals>.compute_sum  s2    LLE*<MND2dJ/N
 	
Ns   9
)r  r      )r  )r   r   )rg   r   rU   )r/   r0   r  r  r  r   is          rZ   simple_sparse_reduce_testsr    s    
?s 
? 
?# 
?s 
?
C 
 H!,H!,H!,H+H+H+
 	 z* :$q(*z*AB @EZ?PQ![Z*45Q	
  Rs$   5CC C/CC
Cc           
      f   t         j                  j                         }t        rt         j                  j                         }t
        rt         j                  j                         }t        |      }d}| |kD  r|| z  }t        |       D ci c]  }|t        |||z  |dz   |z          }}|S c c}w )zMultigpu tests are designed to simulate the multi nodes with multi
    GPUs on each node. Nccl backend requires equal #GPUs in each process.
    On a single node, all visible GPUs are evenly
    divided to subsets, each process only uses a subset.
    r   )	r   r^   r   r   r   r   r_   rU   list)r0   r.   nGPUsvisible_devicesnGPUs_per_processr  rank_to_GPUs          rZ   init_multigpu_helperr    s     JJ##%E		&&(		&&(ElO E!Z/ z" 	
4$5 5QBS8STUUK  	s   B.tmp_dirinit_methodc                    t        j                         at        j                  t        j
                  d<   t	        j                  t        j                  j                  t        j                  d             t	        j                  t        j                  j                  t        j                  d             t        j                  j                  t        j                  d      }t	        j                  |       | | t        j
                  d<   y t        t        j                  j                  |d      z   t        j
                  d<   y )NTEMP_DIRr  test_dirinit_dirINIT_METHODshared_init_file)
tempfileTemporaryDirectoryr  r=   rG   rH   mkdirpathjoinr   )r  init_dir_paths     rZ   initialize_temp_directoriesr  (  s    ))+G$\\BJJzHHRWW\\',,	23HHRWW\\',,
34GGLLz:MHH]$/

=!$/"'',,-3
 %


=!rj   c                  :    t         t         j                          y y r   )r  cleanupri   rj   rZ   cleanup_temp_dirr  9  s     rj   processcompletion_queuern  c                    |dnt        dt        d|dz              }t        j                         }	 	 |j                  |      S # t        j
                  $ r= | j                         s*|j                         rt        d| j                         cY S Y nw xY w|+t        j                         |z
  }||kD  rt        d| d      S )zGet result from the completion_queue associated with process.

    When the process finished without putting a result or the timeout expired an exception instance will be returnedx   
      rn  zExited with zProcess timed out after s)
maxmintimer	  queueEmptyis_aliveemptyRuntimeErrorexitcode)r  r  rn  queue_timeout
start_timeelapseds         rZ   %retrieve_result_from_completion_queuer  >  s     #?CBCA8N0OMJ
	G#'''>>{{ 	G ##%*:*@*@*B#l73C3C2D$EFF	G iikJ.G #&>wiq$IJJ s   A ABBr  r   c            	       6    e Zd ZdZdZdefdZedefd       Zede	fd       Z
d Z	 dded	edd
f fdZd fdZd fdZdefdZddZddZ G d de      Zede	fd       Zede	dededd
fd       Zdedd
fdZddZddZddZedefd       Z xZS )MultiProcessTestCaser   r  r4   c                      yr,  ri   selfs    rZ   _should_stop_test_suitez,MultiProcessTestCase._should_stop_test_suites  s    rj   c                      y)NTri   r  s    rZ   destroy_pg_upon_exitz)MultiProcessTestCase.destroy_pg_upon_exit{  s    rj   c                     t         S r   DEFAULT_WORLD_SIZEr  s    rZ   r0   zMultiProcessTestCase.world_size      !!rj   c                 V    t              fd       }t        j                  ||       S )Nc                 j    | j                   | j                  k(  r| j                         y          y r   )r/   MAIN_PROCESS_RANK_join_processesr  r   s    rZ   r   z1MultiProcessTestCase.join_or_run.<locals>.wrapper  s(    yyD222$$R(rj   r
   types
MethodTyper  r   r   s    ` rZ   join_or_runz MultiProcessTestCase.join_or_run  .    	r	 
	 ..rj   method_name
methodNameNc                     |dk7  r|}t         |   |       	 t        | |      }t        | || j	                  |             y # t
        $ r+}|dk7  rt        d| j                   d|       |Y d }~y d }~ww xY wNrunTestzno such test method in : super__init__r   setattrr  AttributeError
ValueError	__class__r  r  r  r   er  s        rZ   r  zMultiProcessTestCase.__init__      
 "$K%		{+BD+t'7'7';< 	Y& !-dnn-=R
|L '	   (A 	A6!A11A6c                    t         |           i | _        g | _        g | _        | j
                  | _        t        j                  d      5 }|j                  | _
        d d d        i | _        y # 1 sw Y   i | _        y xY w)NFdelete)r  setUpspecial_return_code_checksskip_return_code_checks	processesr  r/   r  NamedTemporaryFiler=   	file_namepid_to_pipe)r  r  r  s     rZ   r  zMultiProcessTestCase.setUp  ss     13' .0$**	((6 	$!VVDN	$ 	$ s   A..A>c                 r    t         |           | j                  D ]  }|j                           g | _        y r   )r  tearDownr  	terminate)r  pr  s     rZ   r  zMultiProcessTestCase.tearDown  s3     	AKKM	 rj   c                 F    | j                         j                  d      d   S r  idr   r  s    rZ   r   z'MultiProcessTestCase._current_test_name  s    wwys#B''rj   c                    g | _         t        t        | j                              D ]  }t        j
                  j                         \  }} || j                  j                  dt        |      z   || j                         | j                  |fdt        | dd      i      }|j                          t        j                  d||j                          || j"                  |j                   <   | j                   j%                  |        y )Nprocess fake_pgF)targetr=   r   r   Started process %s with pid %s)r  rU   rg   r0   r   multiprocessingPiper  _runrF   r   r  r   startloggerinfopidr  r  )r  procr/   parent_conn
child_connr  s         rZ   _start_processesz%MultiProcessTestCase._start_processes  s    #doo./ 	+D&+&;&;&@&@&B#K~~**#d)+++-NN	 wtY>G MMOKK8$L,7DW[[)NN!!'*%	+rj   c                     	 t         j                  j                  d       t         j                  j	                  d      j
                  }| j                  |       y # t        $ r Y Fw xY w)Nspawn)r   r  set_start_methodr  get_contextProcessr  )r  r  s     rZ   _spawn_processesz%MultiProcessTestCase._spawn_processes  s[    	!!227; $$009AAd#	  		s   A 	A('A(c                       e Zd ZdZy)MultiProcessTestCase.Eventr   N)rd   re   rf   GET_TRACEBACKri   rj   rZ   Eventr$    s    rj   r&  r/   c                    t         j                  d|       	 t        j                  j	                  | |g      }| |v r| j
                  rt         j                  d|       y | j                         }t         j                  d||       |t        j                  j                  k(  rt        j                  d      5 }t        j                  |       |j                          |j!                  d       | j#                  |j%                                t         j                  d|       d d d        ||v ry # 1 sw Y   xY w)Nz*Starting event listener thread for rank %sz:Pipe closed for process %s, stopping event listener threadzReceived event %s on process %szr+)moder   zProcess %s sent traceback)r  debugr  
connectionwaitclosedrecvr  r  r&  r%  r  r  faulthandlerdump_tracebackflushseeksendread)parent_pipesignal_piper/   ready_pipeseventtmp_files         rZ   _event_listenerz$MultiProcessTestCase._event_listener  s   A4H)4499;:TUKk)%%LLT #((*=udK066DDD!44$? G8$33H= ( a(#((9$?FG k)5  G Gs   :A,D55D>r   r  c                 T     | |      }||_         ||_        |j                  ||       y r   )r/   r  run_testclsr/   r   r  r4  r   r  s          rZ   r  zMultiProcessTestCase._run  s)     9~	"i-rj   c                 :   t         j                  j                  d      \  }}t        j                  t
        j                  ||| j                  fd      }|j                          t        j                  dk7  r2t        j                  dk7  rt         j                  j                  d       dt        j                  d<   t        j                           	  t#        | |              ||j=                  d        |t?        d      |jA                          |jC                          | jD                  r	 tG        jH                          y y # t$        j&                  $ rR}t(        j+                  d	| j                  ||       t        j,                  t.        d
   j0                         Y d }~d }~wt2        $ r t(        j5                  dt7        j8                         | j                  t
        j:                         |j=                  t7        j8                                t        j,                  t
        j:                         Y Zw xY w# ||j=                  d        |t?        d      |jA                          |jC                          w xY w# t>        tJ        f$ r Y y w xY w)NF)duplexT)r  r   daemonrY  darwinr   TORCH_SHOW_CPP_STACKTRACESz4Process %s skipping test %s for following reason: %sr   z;Caught exception: 
%s exiting process %s with exit code: %sz-Expected event_listener_thread to not be None)&r   r  r  	threadingThreadr  r9  r/   r  r   rZ  _C'_set_print_stack_traces_on_fatal_signalrG   rH   r   set_rng_seedr   r   SkipTestr  r  r   r   rb   r   r   	traceback
format_excTEST_ERROR_EXIT_CODEr2  r   r  closer  rD   destroy_process_groupr  )r  r   r4  signal_recv_pipesignal_send_pipeevent_listener_threadses          rZ   r;  zMultiProcessTestCase.run_test  s!   -2-B-B-G-Gu-G-U** ) 0 0'77/;!

 	##%<<7"s||x'? HH<<TB36

/0!!#	 $GD)$&(  + %%d+$,$%TUU!&&($$ **,	 %9    	6KKF			 HHZ	*4455 		@LLN$$&		$99	 Y1134HH)>>?		@  + %%d+$,$%TUU!&&( #J/ sK   E 2J I AF(#I (BI <I ?I  I AJJJc                    g }t        | j                        D ]h  \  }}|j                  | j                  |j                     }	 |j                  t        j                  j                         |j                  ||f       j |D ]x  \  }}	 |j                  d      rK|j                  rt        j                  d|       ;|j!                         }t        j#                  d||       nt        j#                  d|       z y # t        $ r t        j                  d|       Y w xY w# t        $ r t        j                  d|       Y w xY w)Nz>Encountered error while trying to get traceback for process %src  z5Pipe closed for process %s, cannot retrieve tracebackz)Process %s timed out with traceback: 

%sz6Could not retrieve traceback for timed out process: %s)	enumerater  r  r  r  r2  r  r&  r%  r  ConnectionErrorr  	exceptionpollr,  r  r-  r   )r  pipesr  r  piper/   rI  s          rZ   _get_timedout_process_tracebackz4MultiProcessTestCase._get_timedout_process_tracebackG  s0   #DNN3 
	JAw'''4II288FFGLL!T+
	   	JD$99Q<{{S  ! $		ILLEtY LLPRV!	 ' $$X4 #   Ts*   <D3D' >D'D$#D$'E	E	c                    t        | j                               }t        j                         }d}	 	 t        | j                        D ]w  \  }}|j
                  t        j                  k(  s$t        d| d|j
                   d       t        j                  j                         }|D ]  }|j                           d} n |rnt        d | j                  D              rntt        j                         |z
  }	|	|kD  rA| j                          t        d| d       | j                  D ]  }|j                           nt        j                  d	       #t        j                         |z
  }
| j!                  ||
       | j"                  j%                         D ]  }|j'                           y # | j"                  j%                         D ]  }|j'                           w xY w)
NFTProcess z terminated with exit code z", terminating remaining processes.c              3   8   K   | ]  }|j                   d u  y wr   )r  )r.  r	  s     rZ   r/  z7MultiProcessTestCase._join_processes.<locals>.<genexpr>  s     F!qzz-Fs   zTiming out after z" seconds and killing subprocesses.g?)r  r  r  rS  r  r  r  rK  printr   r  active_childrenr  allrY  sleep_check_return_codesr  r  rL  )r  r   rn  r  subprocess_errorr  r	  r^  acr  elapsed_timerX  s               rZ   r  z$MultiProcessTestCase._join_processeso  s   dggi(YY[
 &	%dnn5 DAq zz%9%N%NN&qc)DQZZLPrs +0*?*?*O*O*Q"1 +BLLN++/( $Ft~~FF))+
2W$88:+G94VW "^^ &&

3= @  99;3L$$R6 ((//1 

((//1 

s   9F. +DF. .1Gc           
         | j                   st        j                  d       y| j                   d   }t        | j                         D cg c]&  \  }}|j                  t
        j                  k(  r||f( }}}|r[d}|D ]I  \  }}| j                  |j                     j                         }	|d| dt
        j                   d|	 dz  }K t        |      t        | j                         D ]#  \  }}|j                  t        d| d	| d
       || j                  v ryt        j                         D ]q  }
|j                  |
j                  k(  st        r1t        j!                  d| j#                         |
j$                          yt'        j(                  |
j$                         d}|| j*                  v r| j*                  |   }| j-                  |j                  |d| d|j                   d|j                          yc c}}w )z
        Checks that the return codes of all spawned processes match, and skips
        tests if they returned a return code indicating a skipping condition.
        z<Note: no subprocesses were spawned, test was likely skipped.Nr    r[  z exited with error code z and exception:

 terminated or timed out after  seconds6Skipping %s on sandcastle for the following reason: %szExpected exit code z	 but got z
 for pid: )r   )r  r  warningrS  r  r  rK  r  r  r-  r  r  r   r  rb   r   r  r  rc   r   rH  r   assertEqual)r  r   rd  first_processr  r	  errored_processesr   r  error_messageskipexpected_return_codes               rZ   ra  z(MultiProcessTestCase._check_return_codes  s#    ~~NNN q) "$..1
1zz1FFF F
 

 E/ 
7 $ 0 0 = B B Dqc!9:N:c:c9d e''4oR9 u%% dnn- 	DAqzz!"qc!@hW 	 ---%%' 	:D%%7 
 KKP	
 "++DLL99	:"  ! 000#'#B#B2#F "" %&:%;9]E[E[D\\fgtgxgxfyz 	 	
g
s   
+Hc                      | j                   dk(  S )Nr   r/   r  s    rZ   rm  zMultiProcessTestCase.is_master  s    yyA~rj   r  r  r4   N)rd   re   rf   r  rK  boolr  propertyr  rg   r0   r  rF   r  r  r  r   r  r"  r   r&  staticmethodr9  classmethodr  r;  rY  r  ra  rm  __classcell__r  s   @rZ   r  r  j  s8   
   d   "C " "/ ?H8;	&$(C (+,$    < ..#&.36.	. .7# 7t 7r&P*XJ
X 4  rj   r  c                   >     e Zd Z fdZd ZdefdZddZd Z xZ	S )DistributedTestBasec                     t         |           t        | j                        t        j
                  d<   | j                          y )Nr   )r  r  rF   r0   rG   rH   r"  r  r  s    rZ   r  zDistributedTestBase.setUp  s/    #&t#7

< rj   c                     	 t         j                  j                          	 t	        j
                  | j                         y # t        $ r Y ,w xY w# t        $ r Y y w xY wr   )r   distributedrM  r   rG   remover  OSErrorr  s    rZ   r  zDistributedTestBase.tearDown  sU    	335	IIdnn%  		  		s"   A A 	AA	AAr4   c                 "    d|v ryd|v ryd|v ryyNr^   r(   r   r]   r_   r&   r$   ri   )r  r1   s     rZ   r.   zDistributedTestBase.backend  s$    Vf_f_rj   c                    || j                   }t        j                  |      j                         }t        j                  j                  | j                  |      }t        j                  j                  | j                  |      || j                  |       d| j                  |      v sd| j                  |      v rt        j                  j                         }|re|j                  }t        j                  | d| j                         }t        j                  |       t        j                  j                  |       nt!        d| j                  |       d      t        j                  j"                  j%                         S )Nr.   r0   r/   r>   r(   r&   rH  EExpected to find an accelerator when initializing process group with  backend, but got None)r0   r   get_device_moduler   r  rE   r  init_process_groupr.   r/   acceleratorcurrent_acceleratorr_  r1   set_default_deviceset_device_indexr  distributed_c10d_get_default_group)r  r1   r0   num_visible_devicesr>   r  device_types          rZ   	create_pgzDistributedTestBase.create_pg  sE   J#55f=JJL!!++DNN<OP,,LL(!	 	- 	
 T\\&))Vt||F7K-K++??AK)..Qtyyk&BC((0!!226:"!\\&122HJ    11DDFFrj   c                     t        j                  |      j                         }t        | j                        D ci c]	  }|||z  g c}S c c}w r   r   r  r   rU   r0   r  r1   r  r  s       rZ   rank_to_devicez"DistributedTestBase.rank_to_device$  G    #55f=JJL6;DOO6LMA++,,MMM   Ar   )
rd   re   rf   r  r  rF   r.   r  r  rz  r{  s   @rZ   r}  r}    s%     
 G2Nrj   r}  subtest_configtest_fntest_kwargsc                    t        |j                               }|D cg c]  }|d   	 }}|D cg c]  }|d   	 }}t        j                  | D ]  }	t	        t        ||	d            }
 | j                  di |
5  t        j                  j                           ||i ||
 t        j                  j                          ddd       t        j                           yc c}w c c}w # 1 sw Y   *xY w)a\  
    Runs a test function given by ``test_fn`` as a subtest according to the
    configurations specified by ``subtest_config``. This amortizes the
    costly setup overhead (including process spawn and initializing the
    process group) over the subtests.

    Args:
        subtest_config (Dict[str, List[Any]]): A mapping from subtest
            keyword argument name to a list of its possible values.
        test_fn (Callable): A callable that runs the actual test.
        test_args: Positional arguments to pass to ``test_fn``.
        test_kwargs: Keyword arguments to pass to ``test_fn``.
    r   r   T)strictNri   )r  items	itertoolsproductdictzipsubTestr   _dynamoresetrD   r  )cls_instr  r  	test_argsr  subtest_config_itemsitemsubtest_config_keyssubtest_config_valuesr  subtest_kwargss              rZ   run_subtestsr  )  s    * 9=^=Q=Q=S8T:N%O$d1g%O%OBV-W$d1g-W-W##%:; c"5vdKLX// 	"MM!Y@+@@MM!	" 	 &P-W	" 	"s   C"C'<AC,,C5	c                  n    	 t        j                  g dd      j                  dk(  S # t        $ r Y yw xY w)a   
    If shell command `fi_info -p efa -t FI_EP_RDM` returns exit code 0 then we assume that the machine has
    Libfabric EFA interfaces and EFA software components installed,
    see https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/efa-start.html.
    )fi_infoz-pefaz-t	FI_EP_RDMF)checkr   )
subprocessrun
returncodeFileNotFoundErrorri   rj   rZ   has_efar  K  sA    NN;5j	
  s   %( 	44c                  "    t               rddgS dS )a  
    If the machine has Libfabric EFA interfaces and EFA software components installed it may cause
    'RuntimeError: In operator() at tensorpipe/common/ibv.h:172 "": Operation not supported' if tensorpipe
    uses InfiniBand transport, so we exclude it from tensorpipe transports,
    see https://github.com/pytorch/pytorch/issues/73885 and https://github.com/pytorch/pytorch/issues/65022
    shmuvN)r  ri   rj   rZ   tp_transportsr  _  s     $IE4=/4/rj   c                 d      t        t        |      S d t                fd       }|S )z+
    Wrapper to use with a test method
    )rn  r0   c                      t               t        j                         }fd fd}g }t               D ]=  }t	        j
                  |||f      }|j                          |j                  |       ? |S )Nc                  >     t         j                  j                  k(  S r   rD   r  _worldworlds   rZ   world_is_validzaspawn_threads_and_init_comms.<locals>._run_test_method_with_multi_threads.<locals>.world_is_validx      D118888rj   c                 ~   t        j                  d| |       	                  rt        j                          y y # t        $ rR}t        j                  j                  | t        j                         f       t        j                  |       Y d }~sd }~ww xY w#         rt        j                          w w xY w)Nthreadedr.   r/   r0   r>   )rD   r  BaseExceptionMultiThreadedTestCaseexception_queueputr   exc_infor#   exception_handlerM  )r/   world_pgr>   excallbackr  r0   s       rZ   workerzYspawn_threads_and_init_comms.<locals>._run_test_method_with_multi_threads.<locals>.worker{  s    ##"*E
1
 "#..0 $ ! %55994:PQ!22  "#..0 $s*   A   	B	ABB BB B<r  r   )r!   rD   	HashStorerU   rC  rD  r  r  )	r0   r  global_storer  threadsr/   tr  r  s	   ``     @@rZ   #_run_test_method_with_multi_threadszIspawn_threads_and_init_comms.<locals>._run_test_method_with_multi_threadst  sq    $&~~'	9	1  *% 	D  dE<5PQAGGINN1	
 rj   c                 X    t         j                  j                  j                  d       	   fd      }t        j                  |       t         j                  j                  j                  d       y # t         j                  j                  j                  d       w xY w)NTc                       g i S r   ri   )r   r   r   r  s   rZ   r  z?spawn_threads_and_init_comms.<locals>.wrapper.<locals>.<lambda>  s    D$?$?$? rj   F)r   rE  _distributed_c10d_set_thread_isolation_moder  _join_threads)r  r   r   r  r  r   r0   s   ``` rZ   r   z-spawn_threads_and_init_comms.<locals>.wrapper  sv     	""==dC	I9?G "//>HH&&AA%HEHH&&AA%Hs   %A> >+B))r   spawn_threads_and_init_commsr
   )r   rn  r0   r   r  s   ` ` @rZ   r  r  i  sD     |('j
 	
> 4[
I 
I Nrj   c                       e Zd ZdZ ej
                         ZdZd Z	 dde	de	ddf fdZ
d	 Zd
 Zd fdZ fdZd Zed        Zd Zed        Zed        Zedefd       Zede	fd       ZddddZddddZ xZS )r  a5  
    Test runner that runs all tests with the in-proc process group using
    multiple threads with the threaded process group.

    Each test spawns world_size threads and run the test method in each thread.

    Difference from regular MultiProcess test runner:
    Must explicitly defines SetUp and call self._spawn_threads() to run the tests.
    Cannot use setUp / tearDown (must use perThreadSetup / perThreadShutdown)
        to set up / tear down each thread when running each test.
    No global state possible
        How bad of a limitation is this?
    r   c                 V    t              fd       }t        j                  ||       S )Nc                     | j                   | j                  k(  r| j                  | j                         y          y r   )r/   MAIN_THREAD_RANKr  r  r  s    rZ   r   z2MultiThreadedTestCase.join_or_run.<locals>.wrapper  s.    yyD111""4<<4rj   r  r  s    ` rZ   r  z!MultiThreadedTestCase.join_or_run  r  rj   r  r  r4   Nc                     |dk7  r|}t         |   |       	 t        | |      }t        | || j	                  |             y # t
        $ r+}|dk7  rt        d| j                   d|       |Y d }~y d }~ww xY wr  r  r  s        rZ   r  zMultiThreadedTestCase.__init__  r  r  c                      y r   ri   r  s    rZ   perThreadSetUpz$MultiThreadedTestCase.perThreadSetUp  s    rj   c                      y r   ri   r  s    rZ   perThreadTearDownz'MultiThreadedTestCase.perThreadTearDown  s    rj   c                 x    t         |           | j                  | _        g | _        dt
        j                  d<   y)z
        setUp only set up things in the main thread, if you want to configure things
        in the spawned threads, use perThreadSetUp
        r   rB  N)r  r  r  r/   r  rG   rH   r  s    rZ   r  zMultiThreadedTestCase.setUp  s1    
 	))	36

/0rj   c                 0    t         |           g | _        y)z
        tearDown only set up things in the main thread, if you want to configure things
        in the spawned threads, use perThreadTearDown
        N)r  r  r  r  s    rZ   r  zMultiThreadedTestCase.tearDown  s    
 	rj   c                    t         j                  j                  j                  d       | j                  }t               t        j                         | j                  _	        fd} |       st        d      t        | j                        D ]e  }t        j                  | j                  j                  ||| j                  f      }|j!                          | j"                  j%                  |       g y)zk
        class method to spawn threads and run test, use this method in the SetUp of your TestCase
        Tc                  >     t         j                  j                  k(  S r   r  r  s   rZ   r  z<MultiThreadedTestCase._spawn_threads.<locals>.world_is_valid  r  rj   zInvalid worldr  N)r   rE  r  r  r   r!   rD   r  r  r  r  rU   r0   rC  rD  r  r  r  r  )r  r   r  r/   r  r  s        @rZ   _spawn_threadsz$MultiThreadedTestCase._spawn_threads  s     	""==dC++	$&&*nn&6#	9 //$//* 	#D  ~~**)T4??1SA GGILL"	#rj   c                     | |      }||_         t        |d      rWt        j                         |_        t
        j                  |j                  _        t
        j                  |j                  _	        |j                  |||       y )N_tls)r/   r   rC  localr  r    
_precision	precision_rel_tolrel_tolrun_test_with_threaded_pg)r=  r   r/   r0   r   r  s         rZ   r  zMultiThreadedTestCase._run  sb    9~	 4 !)DI"*"5"5DII ( 1 1DII&&y$
Crj   c                    t        j                  d||| j                  j                         | j	                          	  t        | |              t        j                          | j                          y# t        $ rN}| j                  j                  |t        j                         f       t        j                  |       Y d}~wd}~ww xY w# t        j                          | j                          w xY w)zd
        Run the current test associated with `test_name` using the threaded process group.
        r  r  N)rD   r  r  r  r  r   r  r  r  r   r  r#   r  rM  r  )r  r   r/   r0   r  s        rZ   r  z/MultiThreadedTestCase.run_test_with_threaded_pg  s     	!..--		
 			%$GD)$& &&(""$  	  $$dCLLN%;<.. 	 &&(""$s*   A5 5	C>ACC CC &C5c           
         t         }	 t        |      D ]f  \  }}|j                  t        d|             |j	                         s2t
        j                  j                  |t        t        d| d      d ff       h t        j                          g }| j                  j                         sF| j                  j                         }|j                  |       | j                  j                         sFt                t        j                   j"                  j%                  d       | j'                  |||       y # t                t        j                   j"                  j%                  d       w xY w)Nr   zRank failed to join in under ri  F)r  rS  r  r  r  r  r  r  TimeoutErrorr#   r  r  r	  r  r"   r   rE  r  r  ra  )r=  r  r   rn  idxthreadfailed_ranksfailures           rZ   r  z#MultiThreadedTestCase._join_threads+  s-   !	I(1 VC7O,??$)99== , ,&CG9H$U!" !%	 ##%L))//1--113##G, ))//1 #$HH&&AA%Hgr: #$HH&&AA%Hs   <D9 B,D9 95E.c                    d}d}|D ]&  \  }}|d   }t        |t        j                        r2t        j	                  d|||       |dk  sDt
        d   j                  }Xt        |t              r)d| d| d	}	t        j                  |	       t        |	      t        |t              rEdj                  t        j                  |       }	t        j                  d
|	|       |d| d|	 dz  }t        |t              st        |j                         t"        u s|dk  s|j                   }) t%        |      dkD  rt        |      |dkD  rqt
        j'                         D ]Y  }
||
j                  k(  st(        r#t        j	                  d||
j*                          y t        j                  |
j*                         y y )Nrf  r   r   z3Thread %s skipping test %s for following reason: %sr   r   zThread rh  z	 seconds
z'Caught exception: 
%s exiting thread %sz exited with exception:
rg  rj  )
isinstancer   rH  r  r  r   rb   r  r   r  r   r  rI  format_exception
SystemExitr_  coderg   r   r  r   rc   )r=  r  rn  r   	error_msg	skip_coder/   r  excr   rp  s              rZ   ra  z)MultiThreadedTestCase._check_return_codesI  s    		* 	)ND(1+C#x001I	 q= *9 5 ? ?IC.v%DWIZXS!"3''C+ggi88(CDGdSwtf,EcU"MM	C,>S(Y] #I+	)0 y>Ay))q="))+ >.$T LL
 &//==> rj   c                     t         S r   r  r  s    rZ   r0   z MultiThreadedTestCase.world_size{  r  rj   c                 F    | j                         j                  d      d   S r  r  r  s    rZ   r   z(MultiThreadedTestCase._current_test_name  s     wwys#B''rj   r   rs  c                J    | j                   |k(  r| j                  |||       yy)z
        The reason why we have this util function instead of
        self.assertEqual is all threads are sharing one CPU RNG
        so the assertion result is only reliable on rank 0
        N)r/   rl  r  r   yr   r/   s        rZ   assertEqualOnRankz'MultiThreadedTestCase.assertEqualOnRank  s'     99Q3' rj   c                H    | j                   |k(  r| j                  ||       y y r   )r/   assertNotEqualr  s        rZ   assertNotEqualOnRankz*MultiThreadedTestCase.assertNotEqualOnRank  s#    991% rj   rt  ru  r   )rd   re   rf   __doc__r  Queuer  r  r  rF   r  r  r  r  r  r  ry  r  r  r  ra  rw  rg   r0   r   r
  r  rz  r{  s   @rZ   r  r    s     "ekkmO/ ?H8;	&	7#. D D%. ; ;: /> />b "C " " (C ( (( (&1 & &rj   r  c                        e Zd Zdeej
                  ej                  f   deddf fdZ	dej                  dej                  fdZ
 xZS )SaveForwardInputsModuleforward_inputscast_forward_inputsr4   Nc                 t    t         |           t        j                  dd      | _        || _        || _        y )Nd   )r  r  nnLinearlr  r  r  r  r  r  s      rZ   r  z SaveForwardInputsModule.__init__  s2    
 	3$,#6 rj   r   c                     || j                   | <   | j                  | j                  r3|j                  | j                  j                  j
                              S |      S r   )r  r  r  toweightdtyper  r   s     rZ   forwardzSaveForwardInputsModule.forward  sI    $%D!vv43K3Kadd466==../SSQRSSrj   rd   re   rf   r  r  Moduler   Tensorrv  r  r  rz  r{  s   @rZ   r  r    sT    7RYY457 "7 
	7T T%,, Trj   r  c                        e Zd Zdeej
                  ej                  f   deddf fdZ	dej                  dej                  fdZ
 xZS )SaveForwardInputsModelr  r  r4   Nc                 t    t         |           t        ||      | _        t        ||      | _        || _        y r   )r  r  r  c1c2r  r  s      rZ   r  zSaveForwardInputsModel.__init__  s6    
 	).:MN).:MN,rj   r   c                 `    || j                   | <   | j                  | j                  |            S r   )r  r'  r&  r  s     rZ   r  zSaveForwardInputsModel.forward  s)    $%D!wwtwwqz""rj   r   r{  s   @rZ   r$  r$    sQ    -RYY45- "- 
	-# #%,, #rj   r$  c              #   X  K   |st         j                  j                  |        t         j                  j                         x}r|j                  nd}|t        j                  |      }dt        j                  d<   dt        j                  d<   |rp|rVt         j                  j                  j                  j                  j                         }t        j                  d|| |       nt        j                  || |       t         j                  j!                          t         j                  j"                  j$                  j'                          	 d  t         j                  j!                          t         j                  j"                  j$                  j'                          |rt        j(                          y y # t         j                  j!                          t         j                  j"                  j$                  j'                          |rt        j(                          w w xY ww)	Ncpurb  MASTER_ADDR6789MASTER_PORTfaker  )r.   r/   r0   )r   r  r  r  r_  rD   get_default_backend_for_devicerG   rH   testing	_internalr  r  	FakeStorer  r  r  utilscountersclearrM  )r/   r0   r.   init_pgr  accr  r>   s           rZ   _dynamo_dist_per_rank_initr8    s     **40 "--AACCSC%  55kB +BJJ} &BJJ}MM++77??IIKE##%	 ##G$:V	MM	MM  &&()$$**,&&(  	$$**,&&( s    EH*F> A(H*>A)H''H*c                   @     e Zd ZdZe fd       Ze fd       Z xZS )#DynamoDistributedSingleProcTestCasez
    Test harness for single-process dynamo distributed tests,
    initializes dist process group.

    Prefer this for simple tests, as it's easier to debug.
    c                    t         |           | j                  j                  t	        j
                  t        j                  ddd             d| _        t        j                  j                         j                  }| d| j                   | _        || j                  v rd n| j                  g| _        t        j                   t        j"                  |      | j                  d       y )Nrb  12355)r+  r-  r   rH  r   )r/   r0   )r  
setUpClass_exit_stackenter_contextr   r  rG   rH   r/   r   r  r  r_  r1   
device_idsrD   r  r/  )r=  r1   r  s     rZ   r=  z.DynamoDistributedSingleProcTestCase.setUpClass  s    %%JJ

#.#*	
 ""668==xq
+
!'3::!5CHH://7chhST	
rj   c                 J    t        j                          t        |           y r   )rD   rM  r  tearDownClassr=  r  s    rZ   rB  z1DynamoDistributedSingleProcTestCase.tearDownClass  s    ""$rj   )rd   re   rf   r  ry  r=  rB  rz  r{  s   @rZ   r:  r:    s0     
 
(    rj   r:  c            	       H    e Zd ZdZedefd       Zededededdfd       Z	y)	"DynamoDistributedMultiProcTestCasea   
    Use this for tests that actually run on multiple GPUs.

    Decorate tests with @skip_if_lt_x_gpu(ngpu)

    Note: MultiProcTestCase spawns processes per test and is slow.
    Prefer MultiThreadedTestCase for most tests. Perhaps use this one
    sparingly for integration tests.
    r4   c                 >    t         j                  j                         S r   )r   r  r   r  s    rZ   r0   z-DynamoDistributedMultiProcTestCase.world_size  s      --//rj   r/   r   r  Nc                     t        j                  t        j                                 | |      }||_        ||_        |j                  ||       y r   )r   
addHandlerloggingNullHandlerr/   r  r;  r<  s          rZ   r  z'DynamoDistributedMultiProcTestCase._run  sB     	W0023 9~	"i-rj   )
rd   re   rf   r  rw  rg   r0   ry  rF   r  ri   rj   rZ   rE  rE    sV     0C 0 0 	.	.#&	.36	.		. 	.rj   rE  c                       e Zd ZU dZdZeed<   dZeed<   dZe	dz  ed<    e
d      Ze
ed	<   d
Zeed<   d
Zeed<   ede	dz  fd       Zede	fd       Zed d       Zed        Zede	ddfd       Zed        Zed!d       Zede	defd       Ze fd       Zed        Ze fd       Zd! fdZd Z	 d"de	de	ddf fdZ xZS )#MultiProcContinuousTestr   r0   r/   N	rdvz_filer  )secondsrn  Fpoison_pill_processes_spawnedr4   c                      y)z
        ProcessGroup backend str.
        To be customized by sub test classes, e.g. "nccl".
        Otherwise we return None -- lazily decided by tensor.
        Nri   )r=  s    rZ   backend_strz#MultiProcContinuousTest.backend_str(       rj   c                 \    t         j                  j                         }|y|j                  S )Nr*  )r   r  r  r_  )r=  curr_devices     rZ   r  z#MultiProcContinuousTest.device_type2  s+    '';;=rj   c                      y)z
        ProcessGroup init options.
        To be customized by sub test classes, e.g. ProcessGroupNCCLOpTest
        Here we return None.
        Nri   )r=  high_priority_streams     rZ   optszMultiProcContinuousTest.opts9  rT  rj   c                 J   |t        d      t        |      t        j                  d<   t	        j
                  ||      }t	        j                  | j                         |||| j                         | j                         t        j                  j                         | _        y )Nz!Expected rdvz_file to not be None
LOCAL_RANK)r.   r0   r/   r>   
pg_optionsrn  )r   rF   rG   rH   rD   rE   r  rS  rY  rn  r  r  rX   )r=  r/   r0   rN  r>   s        rZ   _init_pgz MultiProcContinuousTest._init_pgB  s     !DEE $'t9

< y*5OO%!xxzKK	
 &&99;rj   r  c                     |j                  dd      d   } | |      }| j                  |_        | j                  |_        t        ||      }t	        j
                           |di | y )NrR  r   rO  r   ri   )rsplitr/   r0   r   r   rG  )r=  r  r   r   r  r  s         rZ   _run_test_given_idz*MultiProcContinuousTest._run_test_given_idV  sa     NN3N3B7	9~HH	..$	* 	!!# 	&rj   c                    d}d|cxk  r|k  sn t        d| d|       || _        || _        d }	 | j                  |||       t        j                  d       	 |j                         }
t        j                  d	|
        |
nK|%|j                  t        j                  |             S	 | j!                  |
       |j                  |
       vt        j                  d       |st5        j6                          y y # t        $ rO}t        |dd       t        fdt        j                         D        d       }	|	r|	j                  }n Y d }~d }~ww xY w# t"        $ r}t%        |t              rjt        |dd       t        fd
t        j                         D        d       }	|	r4|j                  t        j                  |	j                               Y d }~d}t'        j(                         }dj+                  t-        j.                  |       }t1        d|       }||_        |j                  |       Y d }~od }~ww xY w)NFr   z*Expected 0 <= rank < world_size, got rank=z, world_size=r  c              3   B   K   | ]  }|j                   k(  s|  y wr   rb   r.  vrb   s     rZ   r/  z7MultiProcContinuousTest._worker_loop.<locals>.<genexpr>y  s     Lq1;;)3KL   zSetup completeTz	Got test c              3   B   K   | ]  }|j                   k(  s|  y wr   rc  rd  s     rZ   r/  z7MultiProcContinuousTest._worker_loop.<locals>.<genexpr>  s     Tq1;;);STrf  rf  zException in worker process:
zTerminating ...)r   r/   r0   r]  r   r   nextr   r  rc   r  r)  r	  r  r   rH  r`  r  r  r   r  r  rI  r  r  	__cause__rD   rM  )r=  r/   r0   rN  
task_queuer  raised_exceptioninit_skip_reasonr  
skip_entryr  r  tb_strenhanced_exrb   s                 @rZ   _worker_loopz$MultiProcContinuousTest._worker_loopf  s*    T&J& <TF-PZ|\  #  	LLz95 	%&  nn&GLL9WI./  + $$X%6%67G%HI2&&w/ $$W- V 	&'
  &&(    		FD1ILJ--/LJ #-#5#5  !		> ! 2b*- 'FD 9I "&TJ$5$5$7T"J "(,,X->->z?Q?Q-RS #' <<>!;!;X!FG*-KF8+TU(*% $$[1112s8   C? /"E ?	EAEE	I#A4I AI  Ic                 J   g | _         g | _        g | _        t        j                  d      5 }|j
                  | _        d d d        	 t        j                  j                  d       t        t        |            D ]	  }t        j                  j                         }t        j                  j                         }t        j                  j                  | j                  dt!        |      z   d||| j                  ||f      }|j#                          | j                   j%                  |       | j                  j%                  |       | j                  j%                  |       t&        j)                  d||j*                          y # 1 sw Y   LxY w# t        $ r Y ;w xY w)NFr  r  r  T)r  r=   r@  r   r  )r  task_queuescompletion_queuesr  r  r=   rN  r   r  r  r  rU   rg   r  r!  rp  rF   r  r  r  r)  r  )r=  r0   r  r/   rj  r  r  s          rZ   r"  z(MultiProcContinuousTest._spawn_processes  sX    "((6 	#!FFCM	#
	!!227;
 #j/* 	ND..446J$44::<++33''#d)+JzCST	 4 G MMOMM  )OO"":.!!(()9:LL94M	N	# 	#  		s   FF F	F"!F"r  c                 F   t        j                  | dd      }t        |t              r't        j                  |       }|j                  |      }n| j                  }|dk(  rAt        j                  |      j                         }|dk(  rt        j                  d| d      |S )z
        Get world_size, handling both class variable and property definitions.
        Properties are instance-level and need special handling in class methods.
        r0   NrM  r   zNo z devices available)inspectgetattr_staticr  rw  object__new__fgetr0   r   r  r   r   rH  )r=  r  world_size_attrtemp_instancer0   s        rZ   _get_world_sizez'MultiProcContinuousTest._get_world_size  s     "00lDIox0 #NN3/M(--m<JJ 00=JJLJQ''#k]:L(MNNrj   c                 "    t         |           y)a  
        Class-scope test fixture. Run once for entire test class, before any test starts.
        Note: Process spawning is deferred to setUp to support instantiate_device_type_tests,
        which calls setUpClass during class creation before any tests run.
        N)r  r=  rC  s    rZ   r=  z"MultiProcContinuousTest.setUpClass  s     	rj   c                 :   | j                   ry| j                  j                  d| j                        }t	        |t
              r|j                  |       }n7t	        |t              r|j                  |       }nt        |      r |       }n|}| j                  |      | _        t        | j                        r| j                         n| j                  }|rt        j                  t        j                  t        j                   t        j"                  d}|j                  |      }|  |       st%        j&                  d| d      t(        j+                  d| j,                   d| j                   d|        | j/                  | j                         d	| _         y)
z
        Lazily spawn worker processes on first test run.
        This supports instantiate_device_type_tests which calls setUpClass during
        class creation (before any tests run), when spawning would be premature.
        Nr  )r(   r$   r   r&   z	Backend 'z' is not availablezTesting class z on  T)rQ  __dict__r	  r  r  ry  __func__rw  ry  r   r|  r0   rS  rD   r  r  r&  r-  r   rH  r  r  rd   r"  )r=  device_type_attrr  r.   backend_checkscheck_fns         rZ   _ensure_processes_spawnedz1MultiProcContinuousTest._ensure_processes_spawned  s]    !!
 <<++M3??K&4*33C8K((3 +//4K&'*,K*K ,,[9 (0'@#//#coo....,,..	N &))'2H#HJ'')G9<N(OPPS\\N$s~~.>a}M	
 	S^^,!%rj   c                    | j                   st        | 	          yt        j	                  d| j
                   d       | j                  D ]  }|j                  d        | j                  D ]  }|j                           	 t        j                  | j                         t        j                  d| j                   d       t        | 	          y# t        $ r Y =w xY w)z
        Class-scope test fixture. Run once for entire test class, after all tests finish.
        Tear down the process group if spawned.
        NzJoining z workerszClass z	 finished)rQ  r  rB  r  r)  r0   rr  r  r  r  rG   r  rN  r  r  rd   )r=  rj  r  r  s      rZ   rB  z%MultiProcContinuousTest.tearDownClass*  s     %%G!#x/x89// 	!JNN4 	! }} 	GLLN		IIcmm$ 	fS\\N)45	  		s   C 	C! C!c                    t         |           | j                  j                          | j                  | _        | j                  j                  r&t        j                  d| j                                t        | j                        D ]M  \  }}t        j                  d| d| j                                 |j                  | j                                O y)z5
        Test fixture. Run before each test.
        zPrevious test failed, skipping zSending Rank r  N)r  r  r  r  r  r/   rP  r   rH  r  rS  rr  r  r)  r  )r  r  rj  r  s      rZ   r  zMultiProcContinuousTest.setUpG  s     	 	002 **	 >>%%##&Edggi[$QRR 't'7'78 	&MAzLL=2dggi[9:NN4779%	&rj   c                 V    t              fd       }t        j                  ||       S )Nc           	         | j                   | j                  k(  rgt        j                  d| j	                                 d }t        t        | j                  | j                              D ]  \  }\  }}t        ||t        | j	                                     }|2t        |t        j                        r|}Ot        |t              rTt        j                  d| d| j	                          d| j                   j"                          d| j                   _        |}|| j	                         k7  rt'        d| d| j	                                t        j                  d	| d
| j	                                  ||y          y )NzWaiting for workers to finish r  zDetected failure from Rank z in: z(, skipping rest of tests in Test class: TzExpected rv == self.id(), got z != zMain proc detected rank z
 finished )r/   r  r  r)  r  rS  r  r  rs  r  r  r  r   rH  r  rk  r  rd   rP  r   )r  deferred_exceptionr  r	  r  rvr   s         rZ   r   z>MultiProcContinuousTest._worker_run_main_wait.<locals>.wrapper]  sz   yyD222=dggi[IJ &*"09(>(>?1 ,A,+ ?+[5KB *5 !"h&7&78-/* !"m49!E$'') MEEI^^E\E\D]_ 6:2-/*  TWWY,<RDTWWYKP  LL21#Z	{K5< &1,, 2 rj   r  r  s    ` rZ   _worker_run_main_waitz-MultiProcContinuousTest._worker_run_main_wait\  s/    	r(	 
(	T ..rj   r  r  c                     |dk7  r|}t         |   |       	 t        | |      }t        | || j	                  |             y # t
        $ r+}|dk7  rt        d| j                   d|       |Y d }~y d }~ww xY wr  )r  r  r   r  r  r  r  r  r  s        rZ   r  z MultiProcContinuousTest.__init__  s    
 "$K%		{+BD+t'A'A"'EF 	Y& !-dnn-=R
|L '	r  )Fru  rt  )rd   re   rf   r  r0   rg   rh   r/   rN  rF   r   rn  rP  rv  rQ  ry  rS  r  rY  r]  r`  rp  r"  r|  r=  r  rB  r  r  r  rz  r{  s   @rZ   rL  rL    s   JD#N IsTz "3/GY/K$$C$J    C       < <&  4   O) O)b N N> # #  .   .& .&`    8&*,/f ?H8;	 rj   rL  c                        e Zd ZU eZeed<   edefd       Z	e
defd       ZdeddfdZe
 fd       Zd fd	Zd
 Z xZS )C10dTorchCommsTestBaser0   r4   c                 "    d| v ryd| v ryd| v ryyr  ri   r   s    rZ   r.   zC10dTorchCommsTestBase.backend  s$    Vf_f_rj   c                 `    | j                   }t        |      r |       }| j                  |      S r   )r  r   r.   )r=  r  s     rZ   rS  z"C10dTorchCommsTestBase.backend_str  s)    ooK %-K{{;''rj   r1   Nc                     t         t        t        t        t        t
        d}| j                  |      }||v r||   s| j                  d| d       y y y )N)r$   r&   r(   rcclr,   r*   ztorchcomms z backend is not available)r%   r'   r)   TORCHCOMM_HAS_RCCLr-   r+   r.   skipTest)r  r1   backend_flagsr@   s       rZ   _skip_if_backend_unavailablez3C10dTorchCommsTestBase._skip_if_backend_unavailable  sY    &&&&((
 ||F+=(|1LMMK~5NOP 2M(rj   c                    dt         j                  j                  _        dt        j
                  d<   t        t                     t        j
                  d<   t        |      t        j
                  d<   t        |      t        j
                  d<   |t        j
                  d<   t        | %  |||       | j                         }d|v sd	|v rt         j                  j                         }|rYt        j                  |j                   d
|       }t        j                  |       t         j                  j!                  |       y t#        d| d      y )NTrb  r+  r-  r8   r9   TORCHCOMM_STORE_PATHr(   r&   rH  r  r  )r   r  configuse_torchcommsrG   rH   rF   r   r  r]  rS  r  r  r1   r_  r  r  r  )r=  r/   r0   rN  r.   r  r1   r  s          rZ   r]  zC10dTorchCommsTestBase._init_pg  s   26  /$/

=!$'(8$9

=!'*4y

#$'*:

#$-6

)*z95//#W' 1++??AK)9)9(:!D6&BC((0!!226:"[\c[ddz{  !2rj   c                     | j                   j                  }t        j                  d| j	                         |       t        |      r |       }| j                  t        |             t        | %          y )Nz&Setting up test: %s on device type: %s)
r  r  r  r)  r  r   r  rF   r  r  )r  r  r  s     rZ   r  zC10dTorchCommsTestBase.setUp  sT    nn00=twwy+VK %-K))#k*:;rj   c                     t        j                  |      j                         }t        | j                        D ci c]	  }|||z  g c}S c c}w r   r  r  s       rZ   r  z%C10dTorchCommsTestBase.rank_to_device  r  r  ru  )rd   re   rf   r  r0   rg   rh   rx  rF   r.   ry  rS  r  r]  r  r  rz  r{  s   @rZ   r  r    sv    (J(3   (C ( (Q3 Q4 Q  *Nrj   r  r   r,  )r   ru  r   )r.  	functoolsru  r  rI  r  r  rG   r  r  r   r  rC  r  rI  r  r   collections.abcr   
contextlibr   dataclassesr   datetimer   enumr   r   r	   r
   ior   typingr   r   unittest.mockr   r   torch._dynamo.test_casetorch.cuda.nccltorch.distributedr  rD   torch.nnr  torch._C._autogradr   torch._C._distributed_c10dr   torch._logging._internalr   rC   r   torch.testing._internalr   $torch.testing._internal.common_utilsr   r   r   r   r   r   r   r   r   r   r   r   r    5torch.testing._internal.distributed.multi_threaded_pgr!   r"   r#   r%   r'   r)   r  r-   r+   rI   _backend_flagis_backend_builtglobalsrF   rg   r1   rL   r[   	getLoggerrd   r  setLevelINFOr0  r   HAS_ACCELERATORra   r   r   r   r   r   r   r   r   r   rv  r   r   r   r   r   r   r  r  r  r  r  r   r#  r'  r4  r8  r@  rB  rh   rE  rT  rK  rW  r[  ra  rr  r  getenvr  r  r  r  r  r  r  r  r  r  r!  r  r  r  r  r^   r   r  r  r}  r  r  r  cacher  r  r  r  r!  r  r$  r8  r  	test_caser:  rE  rL  r  ri   rj   rZ   <module>r     s          	   
       $ % !   , ,  "        ) 7 . C 0            $% ':&&x0#GIe$??
? ? LL	?
 ? ? 
?D 
		8	$  4 E? 3x38z 
8
C x$FG	
 Xb"BC x45 8B => 8B >? 8B >? 8B >? 8B >? 8B >? 8B >? 8B >? HR>?  (267!" hr#LM#$ x
V%* 8B DE+, hr#BC-
4 * * **&6   &+ 83 , s s ("
,+\4
N: $+1$ D 
S
sCx *F Fc F# F$ F" 	a 
 
: O)"))$GOPO!$'  +.'(
T 
IC I 2 2(S (c (s (XS 3 2 /3	$	$t	+ 2
S4Z 
4 
" K""**K++11K 4ZK 		K6 QAuzz'>'>'@ AB8 N5N. 5Npd3i( 
 D   &0 
3E7tl&H l&^Tbii T #RYY #  :?#) #)L  %--*A*A*J*J   F.)< .8Gh GTBN4 BNrj   