Deprecated: imwpcache\f884414bce24ee67f\f73723ec7b1919fa5::__construct(): Implicitly marking parameter $YECBGYFECGEAFWHA as nullable is deprecated, the explicit nullable type must be used instead in /www/wwwroot/www.chuangxiangniao.com/wp-content/plugins/imwpcache-dist/build/f884414bce24ee67ff73723ec7b1919fa5.php on line 2

Deprecated: imwpcache\f884414bce24ee67f\f73723ec7b1919fa5::__construct(): Implicitly marking parameter $BBWFDDBHHYHDXXAB as nullable is deprecated, the explicit nullable type must be used instead in /www/wwwroot/www.chuangxiangniao.com/wp-content/plugins/imwpcache-dist/build/f884414bce24ee67ff73723ec7b1919fa5.php on line 2
PyTorch二分类模型精度计算陷阱解析与跨框架对比实践_创想鸟

PyTorch二分类模型精度计算陷阱解析与跨框架对比实践

PyTorch二分类模型精度计算陷阱解析与跨框架对比实践

本文深入探讨了PyTorch二分类模型在精度计算时可能遇到的常见陷阱,特别是当与TensorFlow的评估结果进行对比时出现的显著差异。通过分析一个具体的案例,文章揭示了PyTorch中一个易被忽视的精度计算错误,并提供了正确的实现方式,旨在帮助开发者避免此类问题,确保模型评估的准确性和一致性。

1. 问题现象:PyTorch与TensorFlow的精度差异

在深度学习模型开发过程中,开发者常会遇到在不同框架下实现相似模型时,评估指标出现显著差异的情况。一个典型的二分类问题中,我们观察到以下现象:使用pytorch实现的模型在测试集上仅获得约2.5%的精度,而结构和配置几乎相同的tensorflow模型却能达到约86%的精度。这种巨大的差异通常不是由模型性能本身引起,而是暗示了其中一个框架的评估逻辑可能存在根本性错误。

2. 模型结构与训练配置概览

为了更好地理解问题,我们首先审视两个框架中模型的结构和训练配置。

2.1 PyTorch模型与训练设置

PyTorch模型是一个简单的多层感知机(MLP),包含两个ReLU激活的隐藏层和一个Sigmoid激活的输出层,适用于二分类任务。

import torchimport torch.nn as nnimport torch.optim as optimfrom torch.utils.data import DataLoader, TensorDatasetfrom sklearn.model_selection import train_test_splitimport pandas as pdimport numpy as np# 假设数据加载和预处理已完成# data = pd.read_csv('your_data.csv')# train, test = train_test_split(data, test_size=0.056, random_state=42)# train_X_np = train[["A","B","C", "D"]].to_numpy()# test_X_np = test[["A","B", "C", "D"]].to_numpy()# train_Y_np = train[["label"]].to_numpy()# test_Y_np = test[["label"]].to_numpy()# train_X = torch.tensor(train_X_np, dtype=torch.float32)# test_X = torch.tensor(test_X_np, dtype=torch.float32)# train_Y = torch.tensor(train_Y_np, dtype=torch.float32)# test_Y = torch.tensor(test_Y_np, dtype=torch.float32)# train_dataset = TensorDataset(train_X, train_Y)# batch_size = 64# train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)class SimpleClassifier(nn.Module):    def __init__(self, input_size, hidden_size1, hidden_size2, output_size):        super(SimpleClassifier, self).__init__()        self.fc1 = nn.Linear(input_size, hidden_size1)        self.relu1 = nn.ReLU()        self.fc2 = nn.Linear(hidden_size1, hidden_size2)        self.relu2 = nn.ReLU()        self.fc3 = nn.Linear(hidden_size2, output_size)        self.sigmoid = nn.Sigmoid()    def forward(self, x):        x = self.relu1(self.fc1(x))        x = self.relu2(self.fc2(x))        x = self.sigmoid(self.fc3(x))        return x# input_size = train_X.shape[1]# hidden_size1 = 64# hidden_size2 = 32# output_size = 1# model = SimpleClassifier(input_size, hidden_size1, hidden_size2, output_size)# criterion = nn.BCELoss()# optimizer = optim.Adam(model.parameters(), lr=0.001)# # 原始PyTorch训练循环中的评估部分(存在错误)# num_epochs = 50# for epoch in range(num_epochs):#     # ... (训练代码略)#     with torch.no_grad():#         model.eval()#         predictions = model(test_X).squeeze()#         predictions_binary = (predictions.round()).float()#         accuracy = torch.sum(predictions_binary == test_Y) / (len(test_Y) * 100) # 错误在此行#         if(epoch%25 == 0):#           print("Epoch " + str(epoch) + " passed. Test accuracy is {:.2f}%".format(accuracy))

PyTorch模型使用nn.BCELoss作为损失函数,optim.Adam作为优化器。问题主要出现在评估阶段的精度计算逻辑。

2.2 TensorFlow模型与训练设置

TensorFlow模型同样使用Keras的Sequential API构建了一个相似的MLP结构。

from tensorflow.keras.models import Sequentialfrom tensorflow.keras.layers import Dense# import numpy as np # 假设 train_X, train_Y, test_X, test_Y 已经准备好为 numpy 数组# # 假设数据加载和预处理已完成# # model_tf = Sequential()# # model_tf.add(Dense(64, input_dim=len(train_X[0]), activation='relu'))# # model_tf.add(Dense(32, activation='relu'))# # model_tf.add(Dense(1, activation='sigmoid'))# # Compile the model# # model_tf.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])# # model_tf.fit(train_X, train_Y, epochs=50, batch_size=64, verbose=0)# # Evaluate the model# # loss_tf, accuracy_tf = model_tf.evaluate(test_X, test_Y, verbose=0)# # print(f"Loss: {loss_tf}, Accuracy: {accuracy_tf}")

TensorFlow模型在编译时直接指定了metrics=[‘accuracy’],这使得其在训练和评估时能够自动计算并报告正确的精度。

通过对比可以看出,两个框架的模型结构、损失函数和优化器选择都非常相似,主要的差异在于PyTorch的精度计算是手动实现,而TensorFlow则使用了内置的可靠指标。

3. PyTorch精度计算的症结所在

问题的核心在于PyTorch评估代码中的精度计算方式。

3.1 错误代码分析

原始PyTorch代码中的精度计算如下:

accuracy = torch.sum(predictions_binary == test_Y) / (len(test_Y) * 100)

让我们逐步分析这行代码:

predictions_binary == test_Y:这是一个布尔张量,表示每个预测是否与真实标签匹配。torch.sum(…):计算布尔张量中 True 的数量,即正确分类的样本数。len(test_Y):获取测试集中的总样本数。(len(test_Y) * 100):这是问题的关键所在。分母被错误地乘以了100。

正确的精度计算逻辑应该是:(正确分类样本数 / 总样本数) * 100%。例如,如果有86个正确预测和100个总样本,实际精度应为 (86 / 100) * 100% = 86%。然而,原始代码的计算是 (86 / (100 * 100)),即 86 / 10000 = 0.0086。如果再将其格式化为百分比,就会显示为 0.86%,或者在某些情况下,如果期望输出的是0-100的数值,则会是 0.86,与86%相去甚远。原始代码中 format(“{:.2f}%”.format(accuracy)) 会将 0.0086 格式化为 0.86%,而不是 86.00%。因此,PyTorch代码中2.5%的低精度实际上是由于计算公式中分母多乘了一个100,导致最终结果被额外缩小了100倍。

3.2 正确的精度计算方法

为了获得正确的百分比精度,我们需要修正计算公式:

# 假设 predictions_binary 是模型输出经过 Sigmoid 后,再四舍五入得到的二值预测 (0或1)# 假设 test_Y 是真实的二值标签 (0或1)# 计算正确预测的数量correct_predictions = (predictions_binary == test_Y).sum().item()# 获取总样本数total_samples = test_Y.size(0) # 或者 len(test_Y)# 计算精度(0-100

以上就是PyTorch二分类模型精度计算陷阱解析与跨框架对比实践的详细内容,更多请关注创想鸟其它相关文章!

版权声明:本文内容由互联网用户自发贡献,该文观点仅代表作者本人。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。
如发现本站有涉嫌抄袭侵权/违法违规的内容, 请发送邮件至 chuangxiangniao@163.com 举报,一经查实,本站将立刻删除。
发布者:程序猿,转转请注明出处:https://www.chuangxiangniao.com/p/1376098.html

赞 (0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
使用 NumPy 计算 3D 数组列均值并填充 NaN 值
上一篇 2025年12月14日 15:36:44
SQLAlchemy异步会话与PostgreSQL连接管理深度解析
下一篇 2025年12月14日 15:36:52

相关推荐

  • 并发处理共享列表并收集结果的方案

    并发处理共享列表并收集结果的方案并发处理共享列表并收集结果的方案并发处理共享列表并收集结果的方案并发处理共享列表并收集结果的方案

    本文旨在介绍如何利用 Java 并行流高效地处理大型列表,尤其是在每个元素的处理过程耗时较长的情况下。并行流能够将列表分割成多个子任务,并在多个线程上并发执行,从而显著提升处理速度。但同时,并发编程也带来了共享资源同步的问题,需要谨慎处理。 使用并行流并发处理列表 假设我们有一个 Foo 类,其 p…

    2026年9月25日 • 用户投稿
    000
  • 参加PHP+MySQL就业培训后能获得的岗位有哪些

    参加php+mysql就业培训后,你可以获得以下岗位:1. web开发工程师,利用php和mysql开发动态网站和web应用程序;2. 后端开发工程师,使用php构建后端服务和api;3. 全栈开发工程师,结合前端技术进行全站开发;4. 数据库管理员,负责mysql数据库的设计、优化和维护;5. 软…

    2026年9月25日
    200
  • 高效并发处理共享列表与结果收集的Java教程

    高效并发处理共享列表与结果收集的Java教程高效并发处理共享列表与结果收集的Java教程高效并发处理共享列表与结果收集的Java教程高效并发处理共享列表与结果收集的Java教程

    本文介绍了如何利用Java并发特性,特别是并行流(Parallel Streams),来高效处理共享列表,并将处理结果进行收集。针对耗时操作,通过将列表分割成子列表,并利用并行流并发执行,可以显著提高处理效率。同时,强调了在并发环境下对共享资源进行同步的重要性,并提供了收集处理结果的示例代码。 在处…

    2026年9月25日 • 用户投稿
    000
  • AI Overviews能否用于电商搜索 产品信息摘要在购物场景下的使用体验

    AI Overviews能否用于电商搜索 产品信息摘要在购物场景下的使用体验AI Overviews能否用于电商搜索 产品信息摘要在购物场景下的使用体验AI Overviews能否用于电商搜索 产品信息摘要在购物场景下的使用体验AI Overviews能否用于电商搜索 产品信息摘要在购物场景下的使用体验

    随着人工智能技术的发展,AI Overviews作为一种通过整合信息提供摘要的搜索功能,正逐渐改变用户获取信息的方式。本文将探讨AI Overviews是否以及如何在电商搜索场景下应用,特别关注产品信息摘要对于用户购物体验的影响。我们将讲解其运作原理、潜在优势、面临挑战以及优化体验的过程,帮助理解这…

    2026年9月25日 • 用户投稿
    000
  • AI 图像水印失守!开源工具 5 分钟内抹除所有水印

    AI 图像水印失守!开源工具 5 分钟内抹除所有水印AI 图像水印失守!开源工具 5 分钟内抹除所有水印AI 图像水印失守!开源工具 5 分钟内抹除所有水印AI 图像水印失守!开源工具 5 分钟内抹除所有水印

    ai 图像的水印技术正面临重大挑战! 一种名为 UnMarker 的新型去水印技术横空出世,宣称可在短短5分钟内清除市面上绝大多数 AI 生成图像中的水印。 该技术已成功完全破解谷歌的 HiDDeN 水印系统,对另一款 Google 水印技术 SynthID 的破解率也达到了79%。 更令人震惊的是…

    2026年9月25日 • 用户投稿
    000
  • AI Overviews与传统摘要工具有何不同 模型机制与结果效果的差异分析

    AI Overviews与传统摘要工具有何不同 模型机制与结果效果的差异分析AI Overviews与传统摘要工具有何不同 模型机制与结果效果的差异分析AI Overviews与传统摘要工具有何不同 模型机制与结果效果的差异分析AI Overviews与传统摘要工具有何不同 模型机制与结果效果的差异分析

    本文将探讨AI Overviews与传统摘要工具之间的核心差异,重点分析它们在模型机制和结果效果上的不同。通过理解这两种技术的底层原理和最终呈现形式,用户可以更好地认识到它们各自的优势和应用场景。文章将分步讲解这些差异点,帮助您掌握如何区分并理解它们的工作方式。 ☞☞☞AI 智能聊天, 问答助手, …

    2026年9月25日 • 用户投稿
    000
  • 如何在微服务之间共享静态数据

    如何在微服务之间共享静态数据如何在微服务之间共享静态数据如何在微服务之间共享静态数据如何在微服务之间共享静态数据

    微服务架构的本质决定了微服务之间无法直接共享静态变量。正如上面摘要所说,每个微服务都是一个独立的进程,拥有自己的内存空间,静态变量只在其所属的进程内有效。试图在一个微服务中访问另一个微服务的静态变量,就像试图在一个独立的Java程序中访问另一个程序的变量一样,是不可能的。 微服务架构的独立性 微服务…

    2026年9月25日 • 用户投稿
    100
  • AI Overviews在多标签页面下怎么使用 页面复杂结构下的信息筛选能力说明

    AI Overviews在多标签页面下怎么使用 页面复杂结构下的信息筛选能力说明AI Overviews在多标签页面下怎么使用 页面复杂结构下的信息筛选能力说明AI Overviews在多标签页面下怎么使用 页面复杂结构下的信息筛选能力说明AI Overviews在多标签页面下怎么使用 页面复杂结构下的信息筛选能力说明

    本文旨在说明AI Overviews如何在处理多标签页面的信息过载以及复杂网页结构的阅读挑战中发挥作用。我们将探讨AI Overviews如何帮助用户快速掌握多个来源或单个冗长页面中的关键信息,通过智能化的方式进行信息筛选和整合,从而提升信息获取的效率。文章将提供一个基本的操作流程说明,方便用户理解…

    2026年9月25日 • 用户投稿
    100
  • [python]windows上通过whl文件安装triton模块

    [python]windows上通过whl文件安装triton模块[python]windows上通过whl文件安装triton模块[python]windows上通过whl文件安装triton模块[python]windows上通过whl文件安装triton模块

    在windows系统中,使用.whl文件安装triton是一个简单且高效的方法。以下是完整的操作流程说明: 一、检查系统配置 Python版本:首先确认已安装Python,并确保其版本与你要安装的Triton .whl 文件兼容。例如,若下载的是triton-2.0.0-cp310-cp310-wi…

    2026年9月25日 • 用户投稿
    300
  • 如何在微服务之间共享静态数据?

    如何在微服务之间共享静态数据?如何在微服务之间共享静态数据?如何在微服务之间共享静态数据?如何在微服务之间共享静态数据?

    在微服务架构中,各个服务都是独立的部署单元,拥有各自的内存空间。如同上述摘要所述,直接通过静态变量在不同的微服务之间共享数据是不可能的。 试图在一个微服务中设置静态变量的值,然后在另一个微服务中访问它,将会得到 null 或初始值,而不是之前设置的值。 这不是 Spring Boot 特有的问题,而…

    2026年9月25日 • 用户投稿
    100
  • Linux系统与Windows系统在资源管理机制上有何差异?

    Linux在服务器领域因cgroups、procfs、ulimit和可调内核参数等机制,提供对资源的精细控制与高透明度;而Windows则通过WDDM、DirectX、优先调度UI线程及完善的驱动生态,优化桌面与多媒体体验,注重流畅性与兼容性。 Linux系统和Windows系统在资源管理机制上存在…

    2026年9月25日
    200
  • 2025年输入指令就可以生成图片的ai免费工具有哪些?

    2025年免费AI图像生成工具将主要来自开源项目、大公司免费额度、独立开发者工具及云平台免费套餐,如Stable Diffusion类开源模型、谷歌微软等集成服务、专注特定领域的在线工具,以及利用AWS、Azure等云平台资源,但通常存在生成速度慢、图像质量低、功能受限、使用次数限制、隐私风险和水印…

    2026年9月25日
    200
  • Micronaut中动态数据结构的类型安全验证策略

    Micronaut中动态数据结构的类型安全验证策略Micronaut中动态数据结构的类型安全验证策略Micronaut中动态数据结构的类型安全验证策略Micronaut中动态数据结构的类型安全验证策略

    本文探讨了在Micronaut应用中,如何有效处理具有动态属性和类型依赖验证的类。通过引入多态接口、特化实现类以及自定义Jackson反序列化器,我们能够实现对复杂动态数据结构的类型安全解析与精细化验证,确保数据完整性和业务规则的正确执行。 动态数据结构的验证挑战 在现代微服务架构中,经常会遇到需要…

    2026年9月25日 • 用户投稿
    1000
  • 从制造到“质造”,格创东智助力TCL摘得中国质量奖

    从制造到“质造”,格创东智助力TCL摘得中国质量奖从制造到“质造”,格创东智助力TCL摘得中国质量奖从制造到“质造”,格创东智助力TCL摘得中国质量奖从制造到“质造”,格创东智助力TCL摘得中国质量奖

    9月16日,tcl科技凭借“极致、领先、协同”的质量管理模式,成功斩获第五届中国质量奖,成为本届广东省及大湾区唯一获此殊荣的企业。这一奖项不仅彰显了tcl在质量管理体系上的卓越成就,也凸显了其智能制造与数字化转型背后的中坚力量——格创东智,在工业质量数智化领域所发挥的关键作用。 作为TCL战略孵化的…

    2026年9月25日 • 用户投稿
    900
  • DeepSeek是否有开源版本 官方提供的开源模型及使用限制说明

    DeepSeek是否有开源版本 官方提供的开源模型及使用限制说明DeepSeek是否有开源版本 官方提供的开源模型及使用限制说明DeepSeek是否有开源版本 官方提供的开源模型及使用限制说明DeepSeek是否有开源版本 官方提供的开源模型及使用限制说明

    对于关注大模型技术的用户而言,了解DeepSeek是否提供开源模型及其相关信息是重要的。DeepSeek确实提供了部分模型作为开源版本,供社区学习和使用。本文旨在详细介绍DeepSeek官方提供的开源模型系列,说明获取这些模型的途径,并重点阐述使用这些开源模型时需要注意的官方限制与许可说明,帮助用户…

    2026年9月25日 • 用户投稿
    000
  • 利用AWS Pinpoint高效发送注册验证码(OTP)教程

    利用AWS Pinpoint高效发送注册验证码(OTP)教程利用AWS Pinpoint高效发送注册验证码(OTP)教程利用AWS Pinpoint高效发送注册验证码(OTP)教程利用AWS Pinpoint高效发送注册验证码(OTP)教程

    本文旨在指导开发者如何高效利用AWS Pinpoint服务发送用户注册验证码(OTP),解决传统AWS SNS在处理动态、未预注册手机号时的局限性。我们将深入探讨Pinpoint作为首选方案的优势,提供具体实现步骤和代码示例,并分享最佳实践,确保OTP消息的可靠、快速送达。 理解注册验证码(OTP)…

    2026年9月25日 • 用户投稿
    100
  • 2025拍照最强的手机排名:最佳夜景拍照手机

    2025拍照最强的手机排名:最佳夜景拍照手机2025拍照最强的手机排名:最佳夜景拍照手机2025拍照最强的手机排名:最佳夜景拍照手机2025拍照最强的手机排名:最佳夜景拍照手机

    随着用户对智能手机摄影性能的要求日益提高,长焦拍摄能力逐渐成为继主摄像头之后影响购机决策的重要因素。尤其是在演唱会、旅行记录、夜间远摄等使用场景中,出色的长焦表现能够显著提升成像清晰度与画面质感。当前市场上,多款旗舰机型在长焦技术方面实现了突破性进展,其中vivo x300 pro、三星galaxy…

    2026年9月25日 • 用户投稿
    000
  • Win10系统下战网无法安装怎么办?

    Win10系统下战网无法安装怎么办?Win10系统下战网无法安装怎么办?Win10系统下战网无法安装怎么办?Win10系统下战网无法安装怎么办?

    战网无法安装怎么处理?当大家遇到战网客户端无法安装的情况时,应该怎么办呢?毕竟组队开黑的小伙伴还在等你。其实,这种现象通常是因为权限问题或是注册表中有之前的残留数据造成的。经过多次尝试,小编终于找到了一个有效的解决方案,接下来就为大家详细讲解具体的操作步骤。 1,按下Ctrl+Alt+Delete组…

    2026年9月25日 • 用户投稿
    100
  • 如何提高debian readdir的并发处理能力

    如何提高debian readdir的并发处理能力如何提高debian readdir的并发处理能力如何提高debian readdir的并发处理能力如何提高debian readdir的并发处理能力

    提升 Debian 系统 readdir 并发处理能力,需要综合考虑文件系统、内核参数、应用程序优化和并行处理技术等多个方面。以下是一些实用建议: 一、选择高效的文件系统 Debian 默认的 ext4/ext3 文件系统性能良好,但对于高并发场景,可以考虑以下选择: XFS: 尤其适用于存储大量文…

    2026年9月25日 • 用户投稿
    100
  • 疑似华为阔比例大折叠曝光:采用7.6-7.7英寸14:10屏幕

    疑似华为阔比例大折叠曝光:采用7.6-7.7英寸14:10屏幕疑似华为阔比例大折叠曝光:采用7.6-7.7英寸14:10屏幕疑似华为阔比例大折叠曝光:采用7.6-7.7英寸14:10屏幕疑似华为阔比例大折叠曝光:采用7.6-7.7英寸14:10屏幕

    9月28日,有数码博主爆料称,疑似华为下一代阔比例大折叠屏手机mate x7正在测试中。该机采用展开后尺寸为7.6-7.7英寸,并采用14:10的比例。该博主称,新机将硬刚苹果折叠屏手机。 华为Mate X6 据CNMO了解,华为Mate X7有望在今年11月份与Mate 80系列一同亮相。在核心性…

    2026年9月25日 • 用户投稿
    100

发表回复

登录后才能评论
关注微信