MCPcopy Create free account
hub / github.com/diviswen/Cycle4Completion / train

Function train

main_code.py:138–257  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

136
137
138def train():
139 with tf.Graph().as_default():
140 with tf.device('/gpu:0'):
141 pointclouds_pl, pointclouds_Y, pointclouds_gt, is_training = MODEL.placeholder_inputs(BATCH_SIZE, NUM_POINT, NUM_POINT_GT)
142
143 batch = tf.get_variable('batch', [], initializer=tf.constant_initializer(0), trainable=False)
144 bn_decay = get_bn_decay(batch)
145 tf.summary.scalar('bn_decay', bn_decay)
146
147 pred_X, pred_Y, pred_Y2X2Y, pred_X2Y2X, X2Y_logits, Y_logits, Y2X_logits, X_logits, X_feats, X2Y_feats,\
148 Y_feats, Y2X_feats, complete_X, incomplete_Y, Y2X2Y_feats, X2Y2X_feats, X2Y_code, Y2X_code, Y2X2Y_code = \
149 MODEL.get_model(pointclouds_pl, pointclouds_Y, is_training, bn_decay, WEIGHT_DECAY)
150 ED_loss, Trans_loss, D_loss, chamfer_loss_X, chamfer_loss_Y, chamfer_loss_X_cycle, chamfer_loss_Y_cycle, D_loss_X, D_loss_Y,\
151 complete_CD, chamfer_loss_partial_X2Y, chamfer_loss_partial_Y2X, code_loss = \
152 MODEL.get_loss(pred_X, pred_Y, pred_Y2X2Y, pred_X2Y2X, X2Y_logits, Y_logits, Y2X_logits, X_logits, X_feats, X2Y_feats, Y_feats, \
153 Y2X_feats, complete_X, incomplete_Y, pointclouds_pl, pointclouds_Y, pointclouds_gt, Y2X2Y_feats, X2Y2X_feats, X2Y_code, Y2X_code, Y2X2Y_code)
154
155 tf.summary.scalar('chamfer_loss_X', chamfer_loss_X)
156 tf.summary.scalar('chamfer_loss_Y', chamfer_loss_Y)
157 tf.summary.scalar('chamfer_loss_X_cycle', chamfer_loss_X_cycle)
158 tf.summary.scalar('chamfer_loss_Y_cycle', chamfer_loss_Y_cycle)
159 tf.summary.scalar('complete_CD', complete_CD)
160 tf.summary.scalar('D_loss_X', D_loss_X)
161 tf.summary.scalar('D_loss_Y', D_loss_Y)
162 tf.summary.scalar('chamfer_loss_partial_X2Y', chamfer_loss_partial_X2Y)
163
164
165 var_list = tf.trainable_variables()
166 ED_var = [var for var in var_list if ('encoder' in var.name) or ('decoder' in var.name)]
167 Trans_var = [var for var in var_list if ('transferer' in var.name)]
168 D_var = [var for var in var_list if 'discriminator' in var.name]
169
170 ED_gradients = tf.gradients(ED_loss, ED_var)
171 Trans_gradients = tf.gradients(Trans_loss, Trans_var)
172 D_gradients = tf.gradients(D_loss, D_var)
173
174 ED_g_and_v = zip(ED_gradients, ED_var)
175 Trans_g_and_v = zip(Trans_gradients, Trans_var)
176 D_g_and_v = zip(D_gradients, D_var)
177
178 learning_rate = get_learning_rate(batch)
179 tf.summary.scalar('learning_rate', learning_rate)
180
181 optimizer = tf.train.AdamOptimizer(learning_rate, beta1=0.9)
182 optimizer_D = tf.train.AdamOptimizer(BASE_LEARNING_RATE, beta1=0.9)
183 optimizer_T = tf.train.AdamOptimizer(BASE_LEARNING_RATE)
184 updata_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS)
185 with tf.control_dependencies(updata_ops):
186 ED_op = optimizer.apply_gradients(ED_g_and_v, global_step=batch)
187 Trans_op = optimizer.apply_gradients(Trans_g_and_v, global_step=batch)
188
189 D_op = optimizer_D.apply_gradients(D_g_and_v, global_step=batch)
190
191
192 saver = tf.train.Saver(max_to_keep=300)
193
194 # Create a session
195 config = tf.ConfigProto()

Callers 1

main_code.pyFile · 0.85

Calls 5

get_bn_decayFunction · 0.85
get_learning_rateFunction · 0.85
log_stringFunction · 0.85
train_one_epochFunction · 0.85
eval_one_epochFunction · 0.85

Tested by

no test coverage detected