损失函数
Last updated
def cmp_cross_entropy_loss(logits, labels, pos_num):
logits_exp_sum = tf.reduce_sum(tf.exp(logits), axis=1)
logits_sum = tf.reduce_sum(tf.multiply(logits, labels), axis=1)
cross_entropy_loss_ = -1.0 * (logits_sum - tf.dtypes.cast(pos_num, tf.float32) * tf.math.log(logits_exp_sum))
cross_entropy_loss = tf.reduce_sum(cross_entropy_loss_)
return cross_entropy_loss, cross_entropy_loss_