TensorFlow代码出现异常类型错误:Maximum操作输入类型不匹配
Let's break down what's causing this error and fix it step by step.
What's Going Wrong?
Your error message flags a type mismatch in the Maximum operation—one input is float32, but the other doesn't match. Here are the key issues in your code:
Mixing Variable Assignment with Tensor Operations
You initializedcurMaxAbsandmaxiastf.Variable, but then overwrote them directly with the output oftf.cond(atf.Tensor). In TensorFlow's graph mode, this doesn't update the original variable—it creates a new tensor with the same name, leading to unexpected type or context issues later on.Potential Type Mismatch in Gradient Tensors
If any gradient tensorgingradsisn'tfloat32,tf.reduce_max(tf.abs(g))will return a tensor of a different type (e.g.,int32ifgis integer-type). Passing this mismatched type totf.maximumalongsidecurMaxAbs(afloat32variable) triggers the error.Python Integer vs. TensorFlow Tensor
When usinglambda: iin yourtf.condformaxi,iis a plain Python integer, not a TensorFlowint32tensor. This can cause implicit conversion bugs that ripple through your computation graph.
Fixed Code Snippet
Here's the corrected version addressing all these issues:
curMaxAbs = tf.Variable(-1.0, tf.float32) maxi = tf.Variable(-1, tf.int32) for i, g in enumerate(grads): # Ensure maxG is float32, even if g has a different type maxG = tf.cast(tf.reduce_max(tf.abs(g)), tf.float32) oCurMaxAbs = curMaxAbs.value() # Get the variable's current value as a tensor # Update curMaxAbs properly using assign() curMaxAbs.assign(tf.cond( tf.greater(maxG, oCurMaxAbs), lambda: maxG, lambda: oCurMaxAbs )) # Update maxi: cast loop index to tf.int32 and use assign() maxi.assign(tf.cond( tf.greater(maxG, oCurMaxAbs), lambda: maxi.value(), lambda: tf.cast(i, tf.int32) )) # This line will now run without type errors if not tf.equal(curMaxAbs, tf.maximum(curMaxAbs, maxG)): # Your logic here pass
Key Fixes Explained
- Use
assign()for Variable Updates: Instead of overwriting variables directly,assign()modifies the originaltf.Variablein place, keeping the tensor context consistent. - Explicit Type Casting:
tf.cast(maxG, tf.float32)ensuresmaxGmatchescurMaxAbs'sfloat32type, even if your gradients are a different type. - Convert Python Integers to Tensors:
tf.cast(i, tf.int32)turns the loop index into a proper TensorFlow tensor, matchingmaxi'sint32type and avoiding conversion bugs. - Retrieve Variable Values Safely:
oCurMaxAbs = curMaxAbs.value()gets the variable's current tensor value, which is more reliable than using the variable directly in conditional operations.
内容的提问来源于stack exchange,提问作者dashnick

