/usr/local/lib64/python3.6/site-packages/torch/nn/parallel/__pycache__
Edit: /usr/local/lib64/python3.6/site-packages/torch/nn/parallel/__pycache__/distributed.cpython-36.pyc (56348B)
3
Eg4 @ sH d dl Z d dlZd dlZd dlZd dlZd dlZd dlmZ d dlZd dl j
Zd dlm
Z
mZ d dlmZmZmZ d dlmZmZ dZej rd dlmZmZ ej
jj rdZd d lmZ d d
lmZ ddl m!Z! d
dl"m#Z# d
dl$m%Z%m&Z&m'Z' dd Z(dd Z)dd Z*dd Z+G dd de
Z,G dd deZ-G dd de!eZ.dS ) N)contextmanager)FunctionVariable)JoinJoinableJoinHook)tree_flattentree_unflattenF)ReduceOp_get_default_groupT)RRef)_get_device_index )Module )_get_stream)gather
is_namedtuplescatter_kwargsc C s: t ot| t}|r$t| j \}}nt| \}}|||fS )N)
RPC_AVAILABLE
isinstancer r local_value)outputoutput_is_rrefoutput_tensor_listtreespec r I/usr/local/lib64/python3.6/site-packages/torch/nn/parallel/distributed.py_tree_flatten_with_rref! s
r c C s t | |} |rt| } | S )N)r r )r r r r r r _tree_unflatten_with_rref, s
r c C st t r"t| tr"| j r"t| j S t| tjr4| gS t| tt frRt
jtt| S t| t
rpt
jtt| j S g S )zI
Recursively find all tensors contained in the specified object.
)r r r Zis_owner
_find_tensorsr torchTensorlisttuple itertoolschainmapdictvalues)objr r r r 3 s
r c / C s ddddddddd d
ddd
ddddddddddddddddddd d!d"d#d$d%d&d'd(d)d*d+d,d-d.d/g/} d0}x4| D ],}|t jkrt j| nd1}|d2||f 7 }qlW t| d S )3NZRANKZ
LOCAL_RANKZ
WORLD_SIZEZMASTER_PORTZMASTER_ADDRZCUDA_VISIBLE_DEVICESZGLOO_SOCKET_IFNAMEZGLOO_DEVICE_TRANSPORTZNCCL_SOCKET_IFNAMEZNCCL_BLOCKING_WAITZ
NCCL_DEBUGZNCCL_DEBUG_SUBSYSZNCCL_IB_DISABLEZNCCL_P2P_DISABLEZNCCL_P2P_LEVELZNCCL_SHM_DISABLEZNCCL_SOCKET_NTHREADSZNCCL_NSOCKS_PERTHREADZ
NCCL_BUFFSIZEZ
NCCL_NTHREADSZ
NCCL_RINGSZNCCL_MAX_NCHANNELSZNCCL_MIN_NCHANNELSZNCCL_CHECKS_DISABLEZNCCL_CHECK_POINTERSZNCCL_LAUNCH_MODEZNCCL_IB_HCAZNCCL_IB_TIMEOUTZNCCL_IB_RETRY_CNTZNCCL_IB_GID_INDEXZ
NCCL_IB_SLZ
NCCL_IB_TCZNCCL_IB_AR_THRESHOLDZNCCL_IB_CUDA_SUPPORTZNCCL_NET_GDR_LEVELZNCCL_NET_GDR_READZNCCL_SINGLE_RING_THRESHOLDZNCCL_LL_THRESHOLDZNCCL_TREE_THRESHOLDZ NCCL_ALGOZ
NCCL_PROTOZNCCL_IGNORE_CPU_AFFINITYZNCCL_DEBUG_FILEZNCCL_COLLNET_ENABLEZNCCL_TOPO_FILEZNCCL_TOPO_DUMP_FILEZNCCL_ASYNC_ERROR_HANDLING zN/Az
env:%s=%s
)osenvironprint)Zrelevant_env_varsZformatted_outputvarvaluer r r _dump_DDP_relevant_env_varsF sh
r1 c @ s$ e Zd Zedd Zedd ZdS )_DDPSinkc G s | j d || _|| _|S )NF)Zset_materialize_gradsreducer
state_dict)ctxr3 r4 inputsr r r forward s
z_DDPSink.forwardc G s6 | j }| j d r.| j d dkr.tjj| jj d|S )Nstatic_graphnum_iterationsr )NN)r4 r Z_execution_engineZqueue_callbackr3 Z_delay_all_reduce)r5 Zgrad_outputsr4 r r r backward s z_DDPSink.backwardN)__name__
__module____qualname__staticmethodr7 r: r r r r r2 s r2 c s2 e Zd Z fddZdd ZedddZ ZS )_DDPJoinHookc s8 t |tstd|jj || _|| j_t j dS )z;
Sets config variables for internal usage.
zQDDP join hook requires passing in a DistributedDataParallel instance as the stateN) r DistributedDataParallelAssertionErrorloggerZ_set_uneven_input_joinddp_divide_by_initial_world_sizesuper__init__)selfrC divide_by_initial_world_size) __class__r r rF s
z_DDPJoinHook.__init__c C sr | j }|jj |j |jdd}|j |j d j dk}||_|sNdS |j |j
rd|j |jj dS )zq
Shadows the DDP collective communication operations in the forward and
backward passes.
T)is_joined_rankr N)
rC r3 _rebuild_buckets_check_and_sync_module_buffers)_check_global_requires_backward_grad_syncwaitresultitemrequire_forward_param_sync_match_all_reduce_for_bwd_passfind_unused_parameters_match_unused_params_allreduceZ_push_all_rebuilt_params)rG rC workZshould_sync_backwardsr r r main_hook s
z_DDPJoinHook.main_hook)is_last_joinerc C s | j j| dS )zj
Syncs the final model to ensure that the model is the same across all
processes.
N)rC _sync_final_model)rG rW r r r post_hook s z_DDPJoinHook.post_hook)r; r<