如何实现check函数:统计两张量中完全相等的子张量数量(多标签分类0/1损失)
解决方案
可以通过以下逻辑实现check函数,统计两个张量中完全匹配的内部子张量数量:
- 逐元素比较两个张量,得到标记元素是否相等的布尔张量
- 对每个内部子张量(行维度)的比较结果取全局逻辑与,判断整行是否完全匹配
- 统计完全匹配的行数,转换为Python整数返回
具体代码实现:
import torch def check(tensor1, tensor2): # 逐元素比较,生成布尔匹配张量 element_match = torch.eq(tensor1, tensor2) # 按行判断是否所有元素都匹配 full_row_match = torch.all(element_match, dim=1) # 统计匹配行数并返回 return int(full_row_match.sum().item())
测试示例:
# 测试第一个案例 a = torch.tensor([[1,0,1,1], [1,0,1,1],[1,0,1,1]]) b = torch.tensor([[0,0,1,1], [1,0,1,1],[1,0,1,0]]) print(check(a,b)) # 输出: 1 # 测试第二个案例 c = torch.tensor([[1,0,1,1], [1,0,1,1],[1,0,1,1]]) d = torch.tensor([[0,0,1,1], [1,0,1,1],[1,0,1,1]]) print(check(c,d)) # 输出: 2
关键函数说明
torch.eq():逐元素对比两个张量,返回同形状的布尔张量,相等位置为True,否则为Falsetorch.all(dim=1):沿行维度(dim=1)对布尔张量取逻辑与,仅当整行所有元素均为True时,该行结果为True.sum().item():统计布尔张量中True的数量(PyTorch中布尔值True对应数值1),并转换为Python原生整数
内容的提问来源于stack exchange,提问作者Saran Pannasuriyaporn
相关产品推荐
相关产品推荐

