/usr/local/lib64/python3.6/site-packages/caffe2/python/helpers/__pycache__
Edit: /usr/local/lib64/python3.6/site-packages/caffe2/python/helpers/__pycache__/train.cpython-36.pyc (1965B)
3
Eg @ sB d dl mZmZ d dlmZ dddZdd Zdd Zd
d ZdS )
)corescope)
caffe2_pb2Nc s> d krt j dkr&| jd d S fdd| jD S d S )N c s g | ]}|j kr|qS )ZGetNameScope).0w) namescoper G/usr/local/lib64/python3.6/site-packages/caffe2/python/helpers/train.py
s z _get_weights..)r ZCurrentNameScopeweights)modelr r )r r
_get_weights s
r c K sP d|kr|d= | j jg |fdgdtjjtjtjdd| | jj ||f|S )N
device_option r )shapevalueZdtyper )
param_init_netConstantFillr ZDataTypeZINT64DeviceOptionr CPUnetZIter)r
blob_outkwargsr r r
iter s r c K s d|kr|d nt j }|d kp*|jtjk}| rd|kr|d dkr| jj|d |d d }| jj|d |d d }| jj||g|fdtj tjdi| n| jj|| d S )Nr Ztop_kr r Z_host)
r ZCurrentDeviceScopeZdevice_typer r r ZCopyGPUToCPUZAccuracyr r )r
Zblob_inr r devZis_cpuZ pred_hostZ
label_hostr r r
accuracy% s
r c C sn |dkrdS | j jg ddg|d}| j jg ddgdd}x0t| D ]$}| j| }| jj||||g| qBW dS )zAdds a decay to weights in the model.
This is a form of L2 regularization.
Args:
weight_decay: strength of the regularization
g Nwdr )r r ONEg ?)r r r Z
param_to_gradr ZWeightedSum)r
Zweight_decayr r paramZgradr r r
add_weight_decay: s
r )N) Z
caffe2.pythonr r Zcaffe2.protor r r r r r r r r
s