Класс CosineSimilarity
Класс CosineSimilarity вычисляет косинусную близость
между метками и предсказаниями. Метрика измеряет
косинус угла между двумя векторами и лежит в диапазоне
от -1 до 1. Значение 1 означает,
что векторы направлены одинаково, 0 - ортогональны,
а -1 - противоположно направлены. Первым параметром
передается имя метрики, вторым - тип данных, третьим -
порог срабатывания. Класс наследуется от tf.keras.metrics.Metric
и подходит для задач, где важно направление вектора,
а не его длина.
Синтаксис
tf.keras.metrics.CosineSimilarity(
name='cosine_similarity',
dtype=None,
axis=-1
)
Пример
Давайте создадим метрику косинусной близости и вычислим ее для двух одинаково направленных векторов:
import tensorflow as tf
metric = tf.keras.metrics.CosineSimilarity()
metric.update_state([1, 2, 3], [2, 4, 6])
print(metric.result().numpy())
Результат выполнения кода:
0.9999999
Пример
Давайте вычислим косинусную близость для ортогональных векторов:
import tensorflow as tf
metric = tf.keras.metrics.CosineSimilarity()
metric.update_state([1, 0], [0, 1])
print(metric.result().numpy())
Результат выполнения кода:
0.0
Пример
Давайте вычислим косинусную близость для противоположно направленных векторов:
Результат выполнения кода:
-0.9999999
Пример
Давайте используем метрику при обучении модели и выведем историю значений:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(3,))
])
model.compile(
optimizer='sgd',
loss='mse',
metrics=[tf.keras.metrics.CosineSimilarity()]
)
x = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=tf.float32)
y = tf.constant([[1], [2], [3]], dtype=tf.float32)
history = model.fit(x, y, epochs=2, verbose=0)
print(history.history['cosine_similarity'])
Результат выполнения кода:
[0.80471295, 0.83201814]
Смотрите также
-
класс
MeanSquaredError,
который вычисляет среднюю квадратичную ошибку -
класс
MeanAbsoluteError,
который вычисляет среднюю абсолютную ошибку -
класс
CosineDecay,
который задает косинусное затухание скорости обучения -
класс
KLDivergence,
который вычисляет дивергенцию Кульбака-Лейблера