ML Community Day is November 9! Join us for updates from TensorFlow, JAX, and more Learn more


Computes the element-wise (weighted) mean of the given tensors.

Inherits From: Metric, Layer, Module

MeanTensor returns a tensor with the same shape of the input tensors. The mean value is updated by keeping local variables total and count. The total tracks the sum of the weighted values, and count stores the sum of the weighted counts.

name (Optional) string name of the metric instance.
dtype (Optional) data type of the metric result.
shape (Optional) A list of integers, a tuple of integers, or a 1-D Tensor of type int32. If not specified, the shape is inferred from the values at the first call of update_state.

Standalone usage:

m = tf.keras.metrics.MeanTensor()
m.update_state([0, 1, 2, 3])
m.update_state([4, 5, 6, 7])
array([2., 3., 4., 5.], dtype=float32)
m.update_state([12, 10, 8, 6], sample_weight= [0, 0.2, 0.5, 1])
array([2.       , 3.6363635, 4.8      , 5.3333335], dtype=float32)
m = tf.keras.metrics.MeanTensor(dtype=tf.float64, shape=(1, 4))
array([[0., 0., 0., 0.]])
m.update_state([[0, 1, 2, 3]])
m.update_state([[4, 5, 6, 7]])
array([[2., 3., 4., 5.]])





View source

Resets all of the metric state variables.

This function is called between epochs/steps, when a metric is evaluated during training.


View source