()
| 136 | |
| 137 | |
| 138 | def 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() |
no test coverage detected