o
    Meq=                    @   s4  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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  m  mZ e dZde_G dd dZG dd	 d	ZG d
d deZ G dd deZ!G dd deZ"G dd deZ#dMddZ$dd Z%dd Z&dd Z'dd Z(dd Z)dd  Z*dNd!d"Z+G d#d$ d$eZ,da-d%d& Z.		dOd'd(Z/d)d* Z0d+d, Z1d-d. Z2d/d0 Z3d1d2 Z4d3d4 Z5d5d6 Z6d7d8 Z7d9d: Z8dPd;d<Z9d=d> Z:d?d@ Z;dAdB Z<dCdD Z=G dEdF dFeZ>dGdH Z?dIdJ Z@dKdL ZAdS )Q    N)closing)	strtoboolrootFc                   @   s   e Zd ZdZdZdZdZdS )DistributeModez\
    There are various mode for fleetrun, each of them is designed for different model.
    r         N)__name__
__module____qualname____doc__Z
COLLECTIVEZPSPS_HETER r   r   UD:\Projects\ConvertPro\env\Lib\site-packages\paddle/distributed/fleet/launch_utils.pyr   )   s
    r   c                   @   s0   e Zd ZdZdZdZdZdZdZdZ	dZdZ
dS )	
DeviceModez
    Training devices type
    r   r   r         N)r   r	   r
   r   UNKNOWNCPUGPUZKUNLUNXPU
ASCEND_NPUMLUr   r   r   r   r   2   s    r   c                   @   sd   e Zd Zdd Zdd Zdd Zdd Zd	d
 Zdd Zdd Z	dd Z
dd Zdd Zdd ZdS )Clusterc                 C   s   d | _ g | _d | _d | _d S N)
job_serverpodshdfsjob_stage_flag)selfr   r   r   r   __init__B      
zCluster.__init__c                 C   s"   d | jdd | jD | j| jS )Nz/job_server:{} pods:{} job_stage_flag:{} hdfs:{}c                 S      g | ]}t |qS r   str).0podr   r   r   
<listcomp>J       z#Cluster.__str__.<locals>.<listcomp>)formatr   r   r   r   r   r   r   r   __str__H   s   zCluster.__str__c                 C   sR   t | jt |jkrdS t| j|jD ]\}}||kr dS q| j|jkr'dS dS NFT)lenr   zipr   )r   clusterabr   r   r   __eq__M   s   zCluster.__eq__c                 C   s   |  | S r   )r2   r   r/   r   r   r   __ne__Z      zCluster.__ne__c                 C   s   t  |j| _d S r   )copyr   r3   r   r   r   update_pods]   s   zCluster.update_podsc                 C   s   t |  S r   )r-   trainers_endpointsr*   r   r   r   trainers_nranks`   r5   zCluster.trainers_nranksc                 C   s
   t | jS r   )r-   r   r*   r   r   r   pods_nranksc      
zCluster.pods_nranksc                 C   s,   g }| j D ]}|jD ]}||j q
q|S r   )r   trainersappendendpoint)r   rr&   tr   r   r   r8   f   s   

zCluster.trainers_endpointsc                 C   s:   g }| j D ]}|jD ]}dd |jD }|| q
q|S )Nc                 S   r"   r   r#   r%   accr   r   r   r'   q   r(   z,Cluster.world_device_ids.<locals>.<listcomp>)r   r<   acceleratorsr=   )r   r?   r&   r@   Zstr_acceleratorsr   r   r   world_device_idsm   s   

zCluster.world_device_idsc                 C   sP   g }| j D ] }d|j|j}|jd kr|jd ks J d||| q|S )Nz{}:{}z{} not a valid endpoint)r   r)   addrportr=   )r   r?   r&   epr   r   r   pods_endpointsu   s   
zCluster.pods_endpointsc                 C   s*   | j D ]}t|t|jkr|  S qd S r   )r   r$   id)r   Zpod_idr&   r   r   r   get_pod_by_id~   s
   
zCluster.get_pod_by_idN)r   r	   r
   r    r+   r2   r4   r7   r9   r:   r8   rD   rH   rJ   r   r   r   r   r   @   s    	r   c                   @   s,   e Zd Zdd Zdd Zdd Zdd Zd	S )
	JobServerc                 C   s
   d | _ d S r   )r>   r*   r   r   r   r       r;   zJobServer.__init__c                 C   s   d | jS )Nz{})r)   r>   r*   r   r   r   r+      r5   zJobServer.__str__c                 C   s   | j |jkS r   )Zendpintr>   r   jr   r   r   r2      r5   zJobServer.__eq__c                 C   
   | |k S r   r   rL   r   r   r   r4      r;   zJobServer.__ne__N)r   r	   r
   r    r+   r2   r4   r   r   r   r   rK      s
    rK   c                   @   s4   e Zd Zdd Zdd Zdd Zdd Zd	d
 ZdS )Trainerc                 C   s   g | _ d | _d | _d | _d S r   )rC   r>   rankstager*   r   r   r   r       r!   zTrainer.__init__c                 C   s   d | j| j| jS )Nz"accelerator:{} endpoint:{} rank:{})r)   rC   r>   rP   r*   r   r   r   r+      s   zTrainer.__str__c                 C   s^   t | jt |jkrdS | j|jks| j|jkrdS t| j|jD ]\}}||kr, dS q!dS r,   )r-   rC   r>   rP   r.   )r   r@   r0   r1   r   r   r   r2      s   zTrainer.__eq__c                 C   rN   r   r   )r   r@   r   r   r   r4      r;   zTrainer.__ne__c                 C      | j S r   rP   r*   r   r   r   rP         zTrainer.rankN)r   r	   r
   r    r+   r2   r4   rP   r   r   r   r   rO      s    rO   c                   @   D   e Zd Zdd Zdd Zdd Zdd Zd	d
 Zdd Zdd Z	dS )Podc                 C   sF   d | _ d | _d | _d | _g | _g | _g | _g | _g | _g | _	d | _
d S r   )rP   rI   rE   rF   r<   serversworkerscoordinatorsheter_workersrC   device_moder*   r   r   r   r       s   
zPod.__init__c                 C   sb   d | j| j| j| j| jdd | jD dd | jD dd | jD dd | j	D dd | j
D 
S )Nzrank:{} id:{} addr:{} port:{} visible_accelerator:{} trainers:{} servers:{}             workers:{} heter_workers:{} coordinators:{}c                 S   r"   r   r#   )r%   r@   r   r   r   r'      r(   zPod.__str__.<locals>.<listcomp>c                 S   r"   r   r#   )r%   sr   r   r   r'      r(   c                 S   r"   r   r#   )r%   wr   r   r   r'      s    c                 S   r"   r   r#   )r%   hr   r   r   r'      r(   c                 S   r"   r   r#   )r%   cr   r   r   r'      r(   )r)   rP   rI   rE   rF   rC   r<   rW   rX   rZ   rY   r*   r   r   r   r+      s   zPod.__str__c                 C   s  | j |j ks| j|jks| j|jks| j|jkr#td| | dS t| jt|jkr:td| j|j dS t	t| jD ]}| j| |j| kr_td| j| |j|   dS qAt| j
t|j
krwtd| j
|j
 dS t	t| j
D ]}| j
| |j
| krtd| j
| |j
|   dS q~t| jt|jkrtd| j|j dS t	t| jD ]}| j| |j| krtd| j| |j|   dS qdS )Nzpod {} != {}Fztrainers {} != {}ztrainer {} != {}zservers {} != {}zworkers {} != {}T)rP   rI   rE   rF   loggerdebugr)   r-   r<   rangerW   rX   )r   r&   ir   r   r   r2      sN   z
Pod.__eq__c                 C   rN   r   r   )r   r&   r   r   r   r4      r;   z
Pod.__ne__c                 C   s   d S r   r   )r   Zres_podsr   r   r   parse_response   s   zPod.parse_responsec                 C   rR   r   rS   r*   r   r   r   rP      rT   zPod.rankc                 C   sD   d}| j D ]	}|d|7 }q|dksJ d| |d d }|S )N z{},z&this pod {} can't see any acceleratorsr   )rC   r)   )r   r?   gr   r   r   get_visible_accelerators   s   
zPod.get_visible_acceleratorsN)
r   r	   r
   r    r+   r2   r4   rd   rP   rg   r   r   r   r   rV      s    	)rV      c                 C   s>   t |}||  t  }t d}|| || |S )Nz>%(levelname)s %(asctime)s %(filename)s:%(lineno)d] %(message)s)logging	getLoggersetLevelStreamHandler	FormattersetFormatter
addHandler)	log_levelnamer`   Zlog_handlerZ
log_formatr   r   r   
get_logger  s   



rr   c                 C   s  t |tu s
J dtd d}d}t| D ]\}}t }	||	_||	_||	_|| }
t|
t|ks5J dt	t|D ]r}t
 }|tjksO|tjksO|tjkrzt|| ttfri|j||  |	j||  n0|j||  |	j||  n|tjkrt|| ttfr|j||  n|j||  d|
|  |_||_|d7 }|	j| q;|j|	 q| |}||j| fS )Ntrainer_endpoints must be listr   r   zMcurrent trainer_endpoints size should be greater equal than acclerators size.%sr   )typelistr   	enumeraterV   rP   rE   r[   r-   rb   rO   r   r   r   r   
isinstancetuplerC   extendr=   r   r>   r<   r   index)node_ipsnode_iptrainer_endpointsr[   devices_per_procr/   Ztrainer_rank	node_rankipr&   cur_node_endpointsrc   trainerpod_rankr   r   r   get_cluster  sB   


r   c                 C   s.  t jdkr4| D ]'}|j d u r.t t |jjtj |j	r$|j	
  td|jj qtd | D ] }|j d u rV|j  |j	rL|j	
  td|jj q6td tddD ]*}d}| D ]}|j d u r{t |jjtj d	}qg|std
  d S td qatd td d S )Nntzterminate process group gid:{}r   zterminate process id:{}r   r   2   FTzterminate all the procszcan't kill all process and exit)osrq   procpollZkillpgZgetpgidpidsignalSIGTERMlog_fncloser`   infor)   timesleep	terminatera   rb   killZSIGKILLfatalexit)procspstepaliver   r   r   terminate_local_procsA  s<   







r   c                  C   s*   zt  } t | }| |fW S    Y d S r   )socketgethostnamegethostbyname)Z	host_namehost_ipr   r   r   get_host_name_ipg  s   

r   c                 K   s6   |t krtn|}|jd|  f|||d d| dS )zAdd argparse's argument.
    Usage:
    .. code-block:: python
        parser = argparse.ArgumentParser()
        add_argument("name", str, "Jonh", "User name.", parser)
        args = parser.parse_args()
    z--z Default: %(default)s.)defaultrv   helpN)boolr   add_argument)argnamerv   r   r   Z	argparserkwargsr   r   r   add_argumentsp  s   
r   c                 C   sZ   dd }t  }d}	 | }||vr|| t|| kr|S |d7 }|dkr,td d S q
)Nc               
   S   sj   t ttjtj!} | tjtjtddd | 	d | 
 d W  d    S 1 s.w   Y  d S )Niir   r   )re   r   )r   r   AF_INETSOCK_STREAM
setsockopt
SOL_SOCKET	SO_LINGERstructpackbindgetsockname)r\   r   r   r   __free_port  s   

$z$find_free_ports.<locals>.__free_portr   Tr   i  z?can't find avilable port and use the specified static port now!)setaddr-   print)numr   Zport_setr   rF   r   r   r   find_free_ports  s    	
r   c                 C   sX   t jdd u rt| }|d urt|}|S tt jd}t|| || |  d}|S )NFLAGS_START_PORTr   )r   environgetr   rw   intrb   )r   offsetports
start_portr   r   r   	get_ports  s   r   c                 C   sF  d}d}d}|   D ]\}}t|t|}q
dd|d| | }dd|| }|| | }	dd	d
g|	  d }
dd	dg|	  d }d	}||
d 7 }|r^|||d |d 7 }n||dd7 }||d 7 }|   D ]'\}}t|trt||krd|dd   }n|}|||d| t|7 }qp||
7 }d|}|S )Nr   (   -   z    z|{{:>{}s}}{}{{:^{}s}}|
 z|{{:>{}s}}{{}}{{:^{}s}}|
z    +re   =+-
r   r   zfleetrun Distributed EnvsValuez... iz
{}
)itemsmaxr-   r)   joinry   r$   )envsheaderspacingZmax_kZmax_vkvZh_formatZl_formatlengthborderlineZdrawsZstr_v_strr   r   r   pretty_print_envs  s4   

r   c                   @   s   e Zd Zdd ZdS )TrainerProcc                 C   s(   d | _ d | _d | _d | _d | _d | _d S r   )r   r   
log_offsetrP   
local_rankcmdr*   r   r   r   r      s   
zTrainerProc.__init__N)r   r	   r
   r    r   r   r   r   r     s    r   c                  G   sH   t | dksJ dt | t | dkr"t| d tsJ | d atS )Nr   zlen(args) {} should <= 1r   )r-   r)   ry   r   _run_with_coverage)argsr   r   r   run_with_coverage  s
   r   c              
   C   s  |d u rt  tj  }nt  |}|dd  |dd  |  }dd |D }g }	t|jD ]\}
}d|j d|j d| 	  d
|  t|
d
dd |jD d
|d	}|d
d d urj|d
 |d
< |dd d urx|d |d< |dd d ur|d |d< t|jdkr|jtjkrdd
dd |jD  |d< n9t|jdkr|jtjkrdd
dd |jD  |d< nt|jdkr|jtjkrdd
dd |jD  |d< t|jdkrdd
dd |jD  |d< tj rt|jdkrdd
dd |jD  |d< || g }t s$tjdddkr(g d}tjdg| |g | }td|| |
dkrZtdt|jt |d td || d }tj!d!krdd ntj"}|d urt#d"| tj$%d#| rt#d$| t&d#| d%}|'d& |'d'
|   W d    n	1 sw   Y  |dd ur|d(( d)krt&d*||
f d+}n	t&d,||
f d+}t)j*|||||d-}nt)j*|||d.}t+ }||_,|j|_|
|_-||_.|r|/ nd |_0||_1|	2| q0|	S )/N
http_proxyhttps_proxyc                 S   s   g | ]}d  |qS :)r   )r%   Zeler   r   r   r'     s    z(start_local_trainers.<locals>.<listcomp>z%dru   ,c                 S   r"   r   r#   rA   r   r   r   r'     r(   )PADDLE_TRAINER_IDZPADDLE_CURRENT_ENDPOINTPADDLE_TRAINERS_NUMPADDLE_TRAINER_ENDPOINTSZPADDLE_RANK_IN_NODEZPADDLE_LOCAL_DEVICE_IDSZPADDLE_WORLD_DEVICE_IDSZPADDLE_CLUSTER_TOPO_PATHPADDLE_RANK_MAPPING_PATHZPADDLE_ENABLE_AUTO_MAPPINGr   c                 S   r"   r   r#   r%   rf   r   r   r   r'     r(   FLAGS_selected_gpusc                 S   r"   r   r#   r   r   r   r   r'   #  r(   ZFLAGS_selected_npusc                 S   r"   r   r#   r   r   r   r   r'   &  r(   ZFLAGS_selected_mlusc                 S   r"   r   r#   r   r   r   r   r'   *  r(   ZFLAGS_selected_acceleratorsc                 S   r"   r   r#   r   r   r   r   r'   .  r(   FLAGS_selected_xpusZWITH_COVERAGEZOFFON)z-mZcoveragerunz--branchz-p-uzstart trainer proc{}  env:{}zYLocal start {} processes. First process distributed environment info (Only For Debug): {}zDistributed Envsr   z~details about PADDLE_TRAINER_ENDPOINTS can be found in {}/endpoints.log, and detail running logs maybe found in {}/workerlog.0r   mkdir -p {}z%s/endpoints.logzrm -f {}/endpoints.logr]   zPADDLE_TRAINER_ENDPOINTS: 
r   ZPADDLE_NEED_RANK_MAPPINGtruez%s/prelaunchlog.%dr0   %s/workerlog.%d)envstdoutstderr
preexec_fn)r   r   )3r6   r   r   poprD   rx   r<   rP   r>   r9   r   r8   r$   rC   r   r-   r[   r   r   r   r   fluidcoreis_compiled_with_xpuupdater   sys
executabler`   ra   r)   r   r   rq   Zsetsidsystempathexistsopenwritelower
subprocessPopenr   r   r   r   tellr   r   r=   )r/   r&   training_scripttraining_script_argslog_dirr   current_envZidsresr   idxr@   proc_envZcoverage_argsr   fnZpre_fnfr   tpr   r   r   start_local_trainers  s   








r  c              
   C   s   | j rIt| j jd5}|| jd |D ]}ztj| W q ty1   tjd| j j  Y qw |	 | _W d    d S 1 sBw   Y  d S d S )Nr?   r   zSUnicodeEncodeError occurs at this line. Please refer to the original log file "%s"
)
r   r   rq   seekr   r   r   r   UnicodeEncodeErrorr  )r  Zfinr   r   r   r   pull_worker_logh  s    "r  c              	   C   s   z?d}g }d}| D ]&}|j r|jdkrt| |j }|d u r#d}q	|dkr/d}||j q	|r=t|  td W |S W |S  t	yR   t
d t|  Y d S  tyf   t
d|| t|      t
d|| t|  Y d S )NFr   Tr   zKeyboardInterrupt, exitzdABORT!!! Out of all {} trainers, the trainer process with rank={} was aborted. Please check its log.)r   r   r  r   r   r=   rP   r   r   KeyboardInterruptr`   warning
SystemExiterrorr)   )r   Znranksr  Z
error_rankr   r   retr   r   r   watch_local_trainersw  sL   


r  c                       | d u rt j }dd td|D }|S td}|d u s"|dkr.dd | dD }|S |d | dD ]}| v sFJ d||f q8 fd	d| dD }td
	| |  |S )Nc                 S   r"   r   r#   r%   xr   r   r   r'     r(   zget_gpus.<locals>.<listcomp>r   CUDA_VISIBLE_DEVICESre   c                 S      g | ]}|  qS r   stripr  r   r   r   r'     r(   r   z4Can't find your gpus %s in CUDA_VISIBLE_DEVICES[%s].c                       g | ]	}  | qS r   r|   r  r  cuda_visible_devices_listr   r   r'         z~Change selected_gpus into reletive values. --ips:{} will change into relative_ips:{} according to your CUDA_VISIBLE_DEVICES:{})
r   r   get_cuda_device_countrb   r   getenvsplitr`   r   r)   )gpusgpus_numZres_gpuscuda_visible_devicesr  r   r!  r   get_gpus  ,   



r*  c                    r  )Nc                 S   r"   r   r#   r  r   r   r   r'     r(   zget_xpus.<locals>.<listcomp>r   XPU_VISIBLE_DEVICESre   c                 S   r  r   r  r  r   r   r   r'     r(   r   z3Can't find your xpus %s in XPU_VISIBLE_DEVICES[%s].c                    r  r   r   r  Zxpu_visible_devices_listr   r   r'     r#  z}Change selected_xpus into reletive values. --ips:{} will change into relative_ips:{} according to your XPU_VISIBLE_DEVICES:{})
r   r   get_xpu_device_countrb   r   r%  r&  r`   r   r)   )xpusZxpus_numZres_xpusZxpu_visible_devicesr  r   r-  r   get_xpus  r+  r0  c                    r  )Nc                 S   r"   r   r#   r  r   r   r   r'     r(   zget_npus.<locals>.<listcomp>r   ZASCEND_VISIBLE_DEVICESre   c                 S   r  r   r  r  r   r   r   r'     r(   r   z6Can't find your npus %s in ASCEND_VISIBLE_DEVICES[%s].c                    r  r   r   r  Znpu_visible_devices_listr   r   r'     r#  zChange selected_npus into reletive values. --ips:{} will change into relative_ips:{} according to your ASCEND_VISIBLE_DEVICES:{})
r   r   get_npu_device_countrb   r   r%  r&  r`   r   r)   )npusZnpus_numZres_npusZnpu_visible_devicesr  r   r1  r   get_npus  r+  r4  c                    r  )Nc                 S   r"   r   r#   r  r   r   r   r'     r(   zget_mlus.<locals>.<listcomp>r   ZMLU_VISIBLE_DEVICESre   c                 S   r  r   r  r  r   r   r   r'     r(   r   z3Can't find your mlus %s in MLU_VISIBLE_DEVICES[%s].c                    r  r   r   r  Zmlu_visible_devices_listr   r   r'     r#  z}Change selected_mlus into reletive values. --ips:{} will change into relative_ips:{} according to your MLU_VISIBLE_DEVICES:{})
r   r   get_mlu_device_countrb   r   r%  r&  r`   r   r)   )mlusZmlus_numZres_mlusZmlu_visible_devicesr  r   r5  r   get_mlus  r+  r8  c                 C   s(  | dkr=t j rt j dkrtd tjS t j r*t j dkr*td tj	S t j
 r=t j dkr=td tjS | dkrOt j dkrOtd tjS | dkrat j dkratd	 tjS | d
krst j dkrstd tj	S | dkrt j dkrtd tjS | dkrtd tjS td)Nheterr   z+launch train in heter mode with GPU device.z+launch train in heter mode with XPU device.z+launch train in heter mode with NPU device.hcclz launch train in ascend npu mode!ncclzlaunch train in GPU mode!bkclzlaunch train in XPU modecnclzlaunch train in MLU modegloozlaunch train in CPU modezDon't supported devices)r   r   is_compiled_with_cudar$  r   r   r   r   r.  r   is_compiled_with_npur2  r   r6  r   r   RuntimeErrorbackendr   r   r   get_device_mode  s<   


rD  c                    s  t | j}g }|tjkrSt| j | jd urMt t| j dks,J d	t | jtt t| j  fddt
jdt D }||fS  }||fS |tjkrt| j| jd urtt| j dksxJ d	t| jttt| j fddt
jdtD }||fS }||fS |tjkrt| j| jd urtt| j dksJ d	t| jttt| j fddt
jdtD }||fS }||fS |tjkr:t| j| jd ur4tt| j dksJ d		t| jttt| j fd
dt
jdtD }||fS }||fS |tjkrmt| drQ| jd u rQt | _| jd u r^dg}||fS dd td| jD }||fS J d	|)Nr   z4gpus' number:{} mod args.nproc_per_node:{} must == 0c                       g | ]
} ||  qS r   r   r%   rc   )r'  nr   r   r'   H      z(get_device_proc_info.<locals>.<listcomp>z4npus' number:{} mod args.nproc_per_node:{} must == 0c                       g | ]
}||   qS r   r   rF  )rG  r3  r   r   r'   T  rH  z4xpus' number:{} mod args.nproc_per_node:{} must == 0c                    rI  r   r   rF  )rG  r/  r   r   r'   `  rH  z4mlus' number:{} mod args.nproc_per_node:{} must == 0c                    rE  r   r   rF  )r7  rG  r   r   r'   l  rH  Zpaddle_cpuonlyc                 S      g | ]}|qS r   r   r  r   r   r   r'   x  s    Fz;Can't support device_mode:{}, support only cpu|gpu|xpu now.)rD  rC  r   r   r*  r'  Znproc_per_noder-   r   r)   sixmovesrb   r   r4  r3  r   r0  r/  r   r8  r7  r   hasattrmultiprocessing	cpu_count)r   r[   r   r   )r'  r7  rG  r3  r/  r   get_device_proc_info;  s   



51


)%




rP  c                 C   s*   t jd| jg| j }t|}|  d S )Nr   )r   r   r  r  r  r  wait)r   r   r   r   r   r   direct_start  s   
rR  c                 C   sn   | dksJ g }|  dD ]"}| dd }| dd }t|| }|d|t|f qd|}|S )zM
    origin_endpoint: ip:port
    user_define_endpoint: ip:(port+offset)
    Nr   r   r   r   )r&  r   r=   r   r$   )Zorigin_endpointsr   Z!paddle_user_define_endpoints_listZip_portr   rF   Znew_portZpaddle_user_define_endpointsr   r   r   get_custom_endpoints  s   
rS  c                 C   s   t |tu s
J d|tjksJ dtd d}t| D ]D\}}t }||_||_||_	|| }	|| }
t
|
dks<J tt
|
D ]}t }d|	|  |_|
| |_|j| qB|j| q| |}||j| fS )Nrs   ,Only support get mapped cluster for gpu now.rt   r   ru   )rv   rw   r   r   r   rx   rV   rP   rE   r[   r-   rb   rO   r>   r<   r=   r   r|   )r}   r~   r   r[   
node_ranksr/   r   r   r&   r   ranks_per_noderc   r   r   r   r   r   'get_mapped_cluster_without_rank_mapping  s*   


rW  c              	      s  |t jks	J dtj }d }t| jd}t|}W d    n1 s&w   Y  g }g }t	|d D ]\}}|
|d  |
|g q5t|dkrR|d }	n| jrY| j}	nt \}
}	|	|v sjJ d|	|f ||	}t|t|ks{J dtd	||	|||  g }g }|D ]] | }tjd
d urttd
d}dd t||t||  D }n)tjdd urttjd}dd t||t||  D }ntt|| }|
 fdd|D  qt||	|||S )NrT  r?   ZmachinesrE   r   r   /Can't find your local ip {%s} in node_ips: {%s}+ranks length should be equal to ips length.Cparsed from args: node_ips:{} node_ip:{} node_rank:{} node_ranks:{}PADDLE_PORTre   c                 S   rJ  r   r   r  r   r   r   r'         zEget_mapped_cluster_from_args_without_rank_mapping.<locals>.<listcomp>r   c                 S   rJ  r   r   r  r   r   r   r'     r\  c                       g | ]}d  |f qS z%s:%dr   r%   rF   r   r   r   r'         )r   r   r   r   r$  r   Zcluster_topo_pathjsonloadrx   r=   r-   hostr   r|   r`   ra   r)   r   r   r   r   r%  rb   r   rW  )r   r[   r(  Zcluster_topo	json_filer}   rU  r	  Zcur_cluster_topor~   _r   
free_portsr   r   r   r`  r   1get_mapped_cluster_from_args_without_rank_mapping  sn   








rh  c                 C   s  t |tu s
J d|tjksJ ddd }td d}t| D ]^\}}	t }
||
_|	|
_||
_	|| }|| }|| }t
t|D ]5}t }|d t||  }t|dks[J d|j||d	  d
||  |_|| |_|
j| qB|j|
 q | |}||j| fS )Nrs   rT  c                 S   sN   t d}|d u s|dkr| S |d}|t| }td| || |S )Nr  re   r   z<Change gpu id from {} to {} based on CUDA_VISIBLE_DEVICES {})r   r%  r&  r|   r$   r`   r   r)   )Zgpu_idr)  r"  Zrelative_idr   r   r   get_relative_gpu_id"  s   


zAget_mapped_cluster_with_rank_mapping.<locals>.get_relative_gpu_idrt   ranksr   z.Only support one process to one device mappingr   ru   )rv   rw   r   r   r   rx   rV   rP   rE   r[   rb   r-   rO   r$   rC   r=   r>   r<   r   r|   )r}   r~   r   r[   rU  node_rank_mappingsri  r/   r   r   r&   r   rV  Zcur_node_rank_mappingrc   r   Zlocal_device_idsr   r   r   r   $get_mapped_cluster_with_rank_mapping  s>   


rl  c              	      s>  |t jks	J dtj }| jptd}d }t|d}t	
|}W d    n1 s-w   Y  dtjd< g }g }g }|D ]$}	||	d  dd t|	d  D }
|
  ||
 ||	 q?t|d	kro|d
 }n| jrv| j}nt \}}||v sJ d||f ||}t|| |ksJ dt|t|ksJ dtd|||||  g }g }|D ]^ | }tjdd urttdd}dd t||t||  D }n*tjdd urttjd}dd t||t||  D }ntt|| }| fdd|D  qt||||||S )NrT  r   r?   re   rE   c                 S   r"   r   r   rF  r   r   r   r'   ]  s    zBget_mapped_cluster_from_args_with_rank_mapping.<locals>.<listcomp>rj  r   r   rX  zGnumber of ranks mapped to one node should not exceed the avaiable ones.rY  rZ  r[  c                 S   rJ  r   r   r  r   r   r   r'     r\  r   c                 S   rJ  r   r   r  r   r   r   r'     r\  c                    r]  r^  r   r_  r`  r   r   r'     ra  )r   r   r   r   r$  rank_mapping_pathr   r%  r   rb  rc  r   r=   rw   keyssortr-   rd  r   r|   r`   ra   r)   r   r   rb   r   rl  )r   r[   r(  rn  Zrank_mappingre  r}   rU  rk  Zcur_rank_mappingZcur_node_rank_listr~   rf  r   rg  r   r   r   r`  r   .get_mapped_cluster_from_args_with_rank_mappingJ  s   











rq  c                   @   rU   )ParameterServerLauncherc                 C   s   || _ || _d| _d| _d| _d| _d| _d| _g | _g | _	d| _
g | _g | _d| _g | _g | _d| _g | _g | _d| _d| _g | _i | _g | _i | _d| _| | d S )NFr   re   T)r   distribute_modewith_coordinator
server_num
worker_numheter_worker_numcoordinator_numserver_endpointsserver_endpoints_ipsserver_endpoints_portworker_endpointsworker_endpoints_ipsworker_endpoints_portheter_worker_endpointsheter_worker_endpoints_ipsheter_worker_endpoints_portcoordinator_endpointscoordinator_endpoints_ipscoordinator_endpoints_portis_localcurrent_node_ipstage_trainer_numstage_heter_map
stage_liststage_device_map	stage_numget_role_endpoints)r   r   rs  r   r   r   r      s6   z ParameterServerLauncher.__init__c              
   C   s&	  |j r;|j | _ |jr)t|jd| j ks$J dt|jd| j |j| _n(t| j d}ddd |D | _n|jdksDJ d|j| _t| jd| _ |jr|j| _|j	rzt|j	d| jksuJ dt|j	d| j|j	| _
nqt| j| j }dd	d |D | _
n^|j	dksJ d
dd |j	dD }t|| _dd |j	dD }d|v rd}t|| j  || j  | j d}g }t| jD ]}|d|| t|| f qd|| _
n|j	| _
|jr/d| _|j| _|jrt|jd| jksJ dt|jd| j|j| _nt| jd}ddd |D | _td | jtjkrY|jdks@J dd| jd< |jd}	tt|	D ]}|	| | j|d < qQ| j
| jd< |jr|jd| _dd | jD | _|jrFt|jdt| jksJ dt|jdt| j|jd}
d| _tt| jD ]}| jdkr|  jd7  _|
| d}t|| j| ksJ d|dd |D }dd |D }d|v rtt|| j| j  | j }g }tt|D ]}|d|| t|| f qd|}nd|}|| j|d < | j|d gt|d  |  j| j| 7  _|  j|7  _qntt| jD ]P}| j| }t|| j | j | j }ddd |D }|| j|d < | j|d gt|d  |  j|7  _| jdkr|  jd7  _|  j|7  _qMn|jdksJ dg | _|jd}
d| _tt|
D ]}|
| d}| jt| dd |D }d d |D }d|v rtt|| j| j  | j }g }tt|D ]}|d|| t|| f qd|}nd|}|| j|d < | j|d gt|d  |  j| jd! 7  _| jdkrB|  jd7  _|  j|7  _q| jg| j | _t| j| _ |j!rb|j!g}ntd| j | j | j }| jdd dd }|d t|d  | _!d"d | jdD | _"d#d | j
dD | _#| jdkrd$d | jdD | _$d%d | jdD | _%d&d | jdD | _&d'd | j
dD | _'g | _(| j"D ]}|| j(vr| j(| q| j#D ]}|| j(vr| j(| q| jtjkr/d(d | jdD | _)d)d | jdD | _*| j)D ]}|| j(vr-| j(| qtt+| j(dkrCd| _,| j(d | _-n0d*| _,t./d+d }|d krXt0 \}| _-n|| _-| jtjkss| j-| j(v ssJ d,| j-| j(f | j-| j(v r| j(1| j-| _2t34d-| j(| j-| j2 d S d S ).Nr   zThe server_num and servers doesn't match. Expect servers endpoints num epual to server_num, but received servers enpoint num: {} and server_num {}r   c                 S      g | ]}d t | qS z
127.0.0.1:r#   r  r   r   r   r'     ra  z>ParameterServerLauncher.get_role_endpoints.<locals>.<listcomp>re   z?The setting of Parameter-Server must has server_num or servers.zThe worker_num and workers doesn't match. Expect workers endpoints num epual to worker_num, but received workers enpoint num: {} and worker_num {}c                 S   r  r  r#   r  r   r   r   r'     ra  z?The setting of Parameter-Server must has worker_num or workers.c                 S      g | ]}|  d d qS r   r   r  r&  r  r   r   r   r'         c                 S      g | ]}t | d qS r   r-   r  r&  r  r   r   r   r'     r  r   i  r   TzThe coordinator_num and coordinators doesn't match. Expect coordinators endpoints num epual to coordinator_num, but received coordinator enpoint num: {} and coordinator_num {}c                 S   r  r  r#   r  r   r   r   r'     ra  z2>>> use default coordinator addr(only one process)zBThe setting of Parameter-Server heter mode must has heter_devices.cpu;r   c                 S   r"   r   rm  )r%   Ztrainer_numr   r   r   r'     s    zThe stage_num and heter_workers doesn't match. Expect heter_workers endpoints stage num epual to heter_worker_num stage, but received heter_workers enpoint stage num: {} and heter_worker_num stage {}z^The heter trainer num in stage {} is not equal in args.heter_worker_num and args.heter_workersc                 S   r  r  r  r  r   r   r   r'   (      c                 S   r  r   r  r  r   r   r   r'   ,  r  c                 S   r  r  r#   r  r   r   r   r'   M  ra  zVThe setting of Parameter-Server heter mode must has heter_worker_num or heter_workers.c                 S   r  r  r  r  r   r   r   r'   _  r  c                 S   r  r   r  r  r   r   r   r'   b  r  r   c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r   r   r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  c                 S   r  r  r  r  r   r   r   r'     r  FPOD_IPzHCan't find your local ip {%s} in args.servers and args.workers ips: {%s}z=parsed from args: node_ips:{} current_node_ip:{} node_rank:{})5ru  rW   r-   r&  r)   ry  r   r   rv  rX   r|  rb   r=   r$   rx  rt  rY   r  r   rs  r   r   Zheter_devicesr  r  rw  Zstage_heter_trainer_numrZ   r  r  r{   r  r  	http_portrz  r}  r  r  r{  r~  r}   r  r  r   r  r  r   r%  r   r|   r   r`   ra   )r   r   r   r}  Zworker_endpoints_lenr   r~  r|  rc   Zheter_devices_listZheter_worker_endpoints_listr  r  Zheter_worker_endpoints_lenr  Znew_heter_worker_endpointsrM   Zip_port_listZheter_trainer_numr  Zhttp_ipr   Zpod_iprf  r   r   r   r    s8  
















*



















z*ParameterServerLauncher.get_role_endpointsc                 C   s  | j | jvrd S td d}d}d}d}d}t| jD ]\}}t }||_||_tt| j	D ]#}	|| j	|	 krQt
 }
d|| j|	 f |
_||
_|d7 }|j|
 q.tt| jD ]&}|| j| krt
 }d|| j| f |_||_d|_|d7 }|j| qYtt| jD ]&}|| j| krt
 }d|| j| f |_||_d|_|d7 }|j| qtt| jD ])}|| j| krt
 }d|| j| f |_||_| j| |_|d7 }|j| q|j| q|j| j }t | _g g g g d| _g g g g d| _ g g g g d| _!| "| j#| | $| j#| | j%r"| &| j#| | j't(j)kr0| *| j#| t+,d-| j#j.| j#j.| j#j.| j#j. t| jd dkr
t| jd D ]"\}	}| jd |	 j/0  t| j!d dkru| j!d |	 1  qTt+,d t| jd	 dkrt| jd	 D ]\}	}| j!d	 |	 1  | jd	 |	 j/2  qt+,d
 t| jd dkrt| jd D ]\}	}| j!d |	 1  | jd |	 j/2  qt+,d t| jd dkr	t| jd D ]\}	}| j!d |	 1  | jd |	 j/2  qt+,d nBt| jd dkr+t| jd D ]\}	}| jd |	 j/0  qt| jd	 dkrLt| jd	 D ]\}	}| jd	 |	 j/0  q<t3j45| jr\t67| j d S d S )Nrt   r   z%s:%sr   )workercoordinatorserverheter_workerzPlease check servers, workers, coordinator and heter_worker logs in {}/workerlog.*, {}/serverlog.* , {}/coordinatorlog.*, and {}/heterlog.*r  zDall workers exit, going to finish parameter server and heter_worker.r  zall heter_worker are killedr  zall parameter server are killedr  zall coordinators are killed)8r  r}   r   rx   rV   rP   rE   rb   r-   rz  rO   r{  r>   rW   r=   r}  r~  rQ   rX   r  r  rY   r  r  r  rZ   r   r   tempfilemkdtempgloo_rendezvous_dirr   cmdslog_fnsstart_pod_serverr   start_pod_workerrt  start_pod_coordinatorrs  r   r   start_pod_heter_workerr`   r   r)   r  r   rQ  r   r   r   r   r   shutilrmtree)r   r/   Zserver_rankZworker_rankZheter_worker_rankZcoordinator_rankr   r   r&   rc   r  rM   r  mr  r   r  r   r   r   r   start_ps  s   






z ParameterServerLauncher.start_psc                 C   s  t j }t|}|dd  |dd  t|jD ]\}}| jtjkrP| j	| j
| j| j|jdd dt| j|jdd tt ddd	| j| jd
}n(| j	| j
| j|jdd dt| j|jdd tt ddd	| j| jd}|| tjd|jg|j }| jd | |dkrtdt|jt|d |j d urt !d|j  t"d|j |f d}	| j#d |	 t$j%|||	|	d}
nt$j%||d}
t& }|
|_'|j(|_(||_)|	|_*|	r|	+ nd |_,||_-| j.d | qd S )Nr   r   r   r   ZPSERVERr   PADDLE_WITH_GLOO03)PADDLE_PSERVERS_IP_PORT_LISTr   PADDLE_COORDINATOR_ENDPOINTS%PADDLE_ALL_HETER_TRAINER_IP_PORT_LISTr[  TRAINING_ROLEr   r  r  PADDLE_GLOO_RENDEZVOUSPADDLE_GLOO_FS_PATHPADDLE_GLOO_HTTP_ENDPOINT)r  r   r  r[  r  r   r  r  r  r  r  r   r  z`Local server start {} processes. First process distributed environment info (Only For Debug): {}r   r   z%s/serverlog.%dr]   r   r   r   r   )/r   r   r6   r   rx   rW   rs  r   r   ry  r|  r  r  r>   r&  r$   rv  r%  r  r  r   r   r   r  r  r  r=   r`   r   r)   r-   r   r  r   r   r  r  r  r   r   rP   r   r   r  r   r   r   )r   r   r&   default_envr  r	  Z
cur_serverr
  r   r  r   r  r   r   r   r  L  s   



z(ParameterServerLauncher.start_pod_serverc              	   C   s0  t j }t|}|dd  |dd  d}g }tj r)t|j}t	|}ntj
 r=tj }dd td|D }t|jD ]R\}}|dkrMdnt|||  }	| jtjkri d| jd| jd	t| jd
| jdt| jdddt| jddd| jd d| jd| jd ddd|jdd d|jdd dt|jdtt dddd| j dd|	|	| j!d}
nOi d| jd| jd	t| jddd
| jd|jdd d|jdd dt|jdtt ddddd | j d!dd"dd#|	d$|	d%| j!}
|"|
 t#j$d&|j%g|j& }| j'd' (| |dkr>t)*d(+t	|jt,|
d) |j-d urit .d*+|j- t/d+|j-|f d,}| j0d' (| t1j2||||d-}nt1j2||d.}t3 }||_4|j|_||_5||_6|r|7 nd |_8||_9| j:d' (| qBd S )/Nr   r   r   c                 S   r"   r   r#   r  r   r   r   r'     r(   z<ParameterServerLauncher.start_pod_worker.<locals>.<listcomp>r  r  r   r   r  PADDLE_STAGE_TRAINERS_NUMSTAGE_ID1	STAGE_NUM*PADDLE_PREVIOUS_HETER_TRAINER_IP_PORT_LISTre   &PADDLE_NEXT_HETER_TRAINER_IP_PORT_LISTr   r  HETER_DEVICE_TYPEr   r  ZTRAINERr  r   r[  r   r  r  r  )r  r   r   r  r,  r  r  r   r   r  r,  r  r   r  z`Local worker start {} processes. First process distributed environment info (Only For Debug): {}r   r   r   r]   r  r  );r   r   r6   r   r   r   r?  r*  r'  r-   r   r.  rb   rx   rX   r$   rs  r   r   ry  r|  rv  r  r  r  r  r  r  r>   r&  rP   r%  r  r  r   r   r   r  r  r  r=   r`   r   r)   r   r  r   r   r  r  r  r   r   r   r   r  r   r   r   )r   r   r&   r  r  heter_device_numdevice_listr	  Z
cur_worker	device_idr
  r   r  r   r  r   r   r   r    s  







	!$
1

	


z(ParameterServerLauncher.start_pod_workerc              	   C   s  t d tj }t|}|dd  |dd  t|jD ]\}}d}i d| jd| jdt	| j
d| jd	t	| jd
dd|jdd d|jdd dt	|jdt	tddddd| jddddd|d|d| j}|| tjd|jg|j }	| jd |	 |dkrtdt|jt|d |jd urt d|j t!d |j|f d!}
| j"d |
 t#j$|	||
|
d"}nt#j$|	|d#}t% }||_&|j|_||_'|
|_(|
r|
) nd |_*|	|_+| j,d | qd S )$Nz">>> entering start_pod_coordinatorr   r   r  r  r   r   r  ZPADDLE_COORDINATOR_NUMr  ZCOORDINATORr  r   r   r[  r   r   r  r  r  r  r   r   r  r,  r  r   r  zeLocal coordinator start {} processes. First process distributed environment info (Only For Debug): {}r   r   z%s/coordinator.%dr]   r  r  )-r   r   r   r6   r   rx   rY   ry  r|  r$   rv  r  rx  r>   r&  rP   r%  r  r  r   r   r   r  r  r  r=   r`   r   r)   r-   r   r  r   r   r  r  r  r   r   r   r   r  r   r   r   )r   r   r&   r  r  r	  Zcur_coordinatorr  r
  r   r  r   r  r   r   r   r    s   




	


z-ParameterServerLauncher.start_pod_coordinatorc              	   C   s  t j }t|}|dd  |dd  d}g }tj r)t|j}t	|}ntj
 r=tj }dd td|D }t|jD ]
\}}|dkrMdnt|||  }	|j}
i d| jd| jd	|
| jd
 krp| j|
d
  ndd| j|
d
  d| jd| j|
 dt|
dt| jd|jdd
 dddt| jdt| jd|jdd dtt ddddd| jddd|	|	| jd}|| tj d|j!g|j" }| j#d $| |dkrt%&d 't	|jt(|d! |j)d ur!t *d"'|j) t+d#|j)|f d$}| j,d $| t-j.||||d%}nt-j.||d&}t/ }||_0|j1|_1||_2||_3|r?|4 nd |_5||_6| j7d $| qBd S )'Nr   r   r   c                 S   r"   r   r#   r  r   r   r   r'   X  r(   zBParameterServerLauncher.start_pod_heter_worker.<locals>.<listcomp>r  r  r   r  r   re   r  r  r  r  r  r[  r   r  ZHETER_TRAINERr   r  r  r  r  r  r  r   )r   r  r,  r  r   r  zfLocal heter_worker start {} processes. First process distributed environment info (Only For Debug): {}r   r   z%s/heterlog.%dr]   r  r  )8r   r   r6   r   r   r   r?  r*  r'  r-   r   r.  rb   rx   rZ   r$   rQ   ry  r|  r  r  r  r  r>   r&  rv  r  r%  r  r  r   r   r   r  r  r  r=   r`   r   r)   r   r  r   r   r  r  r  r   r   rP   r   r   r  r   r   r   )r   r   r&   r  r  r  r  r	  Zcur_heter_workerr  Zstage_idr
  r   r  r   r  r   r   r   r  K  s   








 "%
-z.ParameterServerLauncher.start_pod_heter_workerN)
r   r	   r
   r    r  r  r  r  r  r  r   r   r   r   rr    s    $   Gy?rr  c                 C   s   | dvr
t d|  | dkrtj st d| dkr$tj s$t d| dkr1tj s1t d| d	kr>tj s@t d
d S d S )N)r;  r>  r<  r=  autor:  r9  Zxcclzpaddle.distributed initialize error, backend argument can only be one of 'nccl', 'gloo', 'bkcl', 'auto', 'hccl', 'heter', 'xccl' but got %sr;  zlpaddle.distributed initialize error, your paddle is not compiled with cuda but you assign 'nccl' as backend.r<  zkpaddle.distributed initialize error, your paddle is not compiled with xpu but you assign 'bkcl' as backend.r:  zkpaddle.distributed initialize error, your paddle is not compiled with npu but you assign 'hccl' as backend.r=  zkpaddle.distributed initialize error, your paddle is not compiled with mlu but you assign 'cncl' as backend.)
ValueErrorr   r   r?  r   r@  is_compiled_with_mlurB  r   r   r   check_backend  s.   r  c                 C   s2   | dkrd S t jdrtdt jrtdd S )Nr>  darwinzDYou are going to using gloo on macos, but currently is not supportedzFYou are going to using gloo on windows, but currently is not supported)utilsZOS_NAME
startswithr  Z
IS_WINDOWSrB  r   r   r   block_windows_and_macos  s   r  c                   C   s<   t j rdS t j rdS t j rdS t j rdS dS )Nr;  r<  r:  r=  r>  )r   r   r?  r   r@  r  r   r   r   r   get_backend_by_compile_flag  s   



r  )rh   r   r   )NN)r   )Bri   r   r   r   r6   r   r  r  r  
contextlibr   rN  r   warningsrK  r   rb  ZpaddleZpaddle.fluidr   Zdistutils.utilr   Z*paddle.utils.cpp_extension.extension_utilsr  Zcpp_extensionZextension_utilsrj   r`   	propagater   r   objectr   rK   rO   rV   rr   r   r   r   r   r   r   r   r   r   r   r  r  r  r*  r0  r4  r8  rD  rP  rR  rS  rW  rh  rl  rq  rr  r  r  r  r   r   r   r   <module>   s   
	F!
U)&	
(
)'E
	8>/J      !#