iis服务器助手广告广告
返回顶部
首页 > 资讯 > 后端开发 > Python >Pytorch计算余弦相似度距离——torch.nn.CosineSimilarity函数中的dim参数使用方法
  • 900
分享到

Pytorch计算余弦相似度距离——torch.nn.CosineSimilarity函数中的dim参数使用方法

python机器学习pandas 2023-09-05 10:09:09 900人浏览 八月长安

Python 官方文档:入门教程 => 点击学习

摘要

前言 一、官方函数用法 二、实验验证 1.计算高维数组中各个像素位置的余弦距离 2.验证高维数组中任意一个像素位置的余弦距离 总结 前言 现在要使用PyTorch中自带的torch.nn.CosineSimilarity函数

前言

一、官方函数用法

二、实验验证

1.计算高维数组中各个像素位置的余弦距离

2.验证高维数组中任意一个像素位置的余弦距离

总结


前言

现在要使用PyTorch中自带的torch.nn.CosineSimilarity函数计算两个高维特征图(B,C,H,W)中各个像素位置的特征相似度,即特征图中的每个像素位置上的一个(B,C,1,1)的向量为该位置的特征,总共有BxHxW个特征。

一、官方函数用法

        意思是 dim参数指定了函数在哪个维度上进行余弦距离计算,计算之后该维度会消失,而其他维度的形状保持不变。但是现有的大多数博客将dim的用法复杂化,因此这里进行简单的实验验证,来验证一下上述说法。

二、实验验证

1.计算高维数组中各个像素位置的余弦距离

创造高维数组,在通道维度(即dim=1)上进行向量的余弦距离计算,并查看其中第一批数据中的位置(0,0)上的两个向量之间的余弦距离:

>>> import torch>>> import torch.nn as nn>>> cos = nn.CosineSimilarity(dim=1, eps=1e-6)>>> input1 = torch.randn(3, 64, 100, 128)>>> input2 = torch.randn(3, 64, 100, 128)>>> output = cos(input1, input2)>>> output[0, 0, 0]tensor(-0.1095)

2.验证高维数组中任意一个像素位置的余弦距离

将上述高维数组中的第一批数据中的位置(0,0)上的各个通道数值组成该位置上的特征向量,并计算两个向量间的余弦距离:

>>> import torch>>> import torch.nn as nn>>> cos2 = nn.CosineSimilarity(dim=0, eps=1e-6)>>> input3=input1[0, :, 0, 0]>>> input4=input2[0, :, 0, 0]>>> output2 = cos2(input3, input4)>>> output2tensor(-0.1095)

发现两个距离是相同的,因此dim参数指定了函数在哪个维度上进行余弦距离计算,计算之后该维度会消失,而其他维度的形状保持不变。


总结

  Pytorch中自带的torch.nn.CosineSimilarity函数计算两个高维特征图中各个像素位置的特征相似度,其中dim参数指定了函数在哪个维度上进行余弦距离计算,计算之后该维度会消失,而其他维度的形状保持不变。

来源地址:https://blog.csdn.net/fx714848657/article/details/127384885

--结束END--

本文标题: Pytorch计算余弦相似度距离——torch.nn.CosineSimilarity函数中的dim参数使用方法

本文链接: https://www.lsjlt.com/news/395120.html(转载时请注明来源链接)

有问题或投稿请发送至: 邮箱/279061341@qq.com    QQ/279061341

本篇文章演示代码以及资料文档资料下载

下载Word文档到电脑,方便收藏和打印~

下载Word文档
猜你喜欢
软考高级职称资格查询
编程网,编程工程师的家园,是目前国内优秀的开源技术社区之一,形成了由开源软件库、代码分享、资讯、协作翻译、讨论区和博客等几大频道内容,为IT开发者提供了一个发现、使用、并交流开源技术的平台。
  • 官方手机版

  • 微信公众号

  • 商务合作