PyTorch二分类模型准确率计算陷阱与修正:对比TensorFlow实践

PyTorch二分类模型准确率计算陷阱与修正:对比TensorFlow实践

本文旨在解决PyTorch二分类模型训练过程中,准确率计算可能出现的常见错误,导致结果远低于预期。通过对比TensorFlow的实现,我们将深入分析PyTorch代码中准确率计算的陷阱,并提供正确的计算公式与实践方法,确保模型性能评估的准确性。

1. 问题背景与现象分析

在深度学习二分类任务中,模型性能通常通过准确率(accuracy)来衡量。然而,开发者在使用不同深度学习框架(如pytorch和tensorflow)实现相同模型时,可能会遇到准确率计算结果显著不同的情况。一个常见的问题是,pytorch代码计算出的准确率远低于预期,而tensorflow则表现正常。这往往不是模型本身的差异,而是准确率计算逻辑上的细微错误。

例如,在以下PyTorch二分类模型评估代码中,可能会出现准确率仅为2.5%的异常情况:

# 原始PyTorch准确率计算片段# ...with torch.no_grad():    model.eval()    predictions = model(test_X).squeeze() # 模型输出经过Sigmoid,范围在0-1之间    predictions_binary = (predictions.round()).float() # 四舍五入到0或1    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))# ...

而使用等效的TensorFlow代码,通常能得到合理的准确率(例如86%):

# TensorFlow模型训练与评估片段# ...model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])model.fit(train_X, train_Y, epochs=50, batch_size=64)loss, accuracy = model.evaluate(test_X, test_Y)print(f"Loss: {loss}, Accuracy: {accuracy}")# ...

这种差异的核心原因在于PyTorch代码中准确率计算公式的误用。

2. PyTorch准确率计算错误剖析

上述PyTorch代码中的准确率计算错误主要体现在以下一行:

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

具体分析如下:

除法顺序错误:

为了得到百分比形式的准确率,正确的计算流程应该是:(正确预测数 / 总样本数) * 100。然而,原始代码中的 /(len(test_Y) * 100) 实际上是将正确预测数除以 (总样本数 * 100),这导致结果被额外除以了100,从而使得准确率数值变得非常小(例如,86%的准确率会变成0.86%)。

torch.sum() 返回张量:

torch.sum(predictions_binary == test_Y) 返回的是一个包含正确预测数量的张量(tensor),而不是一个标量(scalar)。虽然PyTorch在某些情况下可以自动进行类型转换,但为了代码的健壮性和清晰性,通常建议使用 .item() 方法将其转换为Python数值类型,尤其是在进行标量运算时。

3. PyTorch中二分类准确率的正确计算方法

要修正PyTorch中的准确率计算,我们需要调整公式以确保正确的百分比转换,并处理好张量到标量的转换。

修正后的准确率计算代码:

# 修正后的PyTorch准确率计算片段# ...with torch.no_grad():    model.eval()    # 确保模型输出和标签形状一致,这里假设test_Y是(N, 1)或(N,)    # 如果model(test_X)输出是(N, 1),则不需要.squeeze()    # 如果model(test_X)输出是(N, 1)且test_Y是(N,),则需要.squeeze()其中一个    # 这里我们假设test_Y是(N, 1),模型输出也是(N, 1),因此不使用.squeeze()    predictions = model(test_X) # 保持(N, 1)形状    predictions_binary = (predictions.round()).float() # 四舍五入到0或1,保持(N, 1)形状    # 计算正确预测的数量    correct_predictions = torch.sum(predictions_binary == test_Y).item()    # 获取总样本数    total_samples = test_Y.size(0) # 等同于 len(test_Y)    # 计算准确率百分比    accuracy = (correct_predictions / total_samples) * 100    if(epoch%25 == 0):      print("Epoch " + str(epoch) + " passed. Test accuracy is {:.2f}%".format(accuracy))# ...

关键修正点:

torch.sum(…).item():将布尔张量的求和结果(正确预测数)转换为Python标量。/ total_samples:计算正确预测的比例。* 100:将比例转换为百分比。

4. 完整的PyTorch二分类模型训练与评估示例

以下是一个集成了正确准确率计算的完整PyTorch二分类模型训练与评估示例:

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# 1. 数据准备 (模拟数据)# 假设你的数据加载和预处理如下:# data = pd.read_csv('your_data.csv')# data['label'] = (data['some_feature'] > threshold).astype(int) # 示例标签生成# ...# 这里使用模拟数据以确保代码可运行np.random.seed(42)num_samples = 1000data = pd.DataFrame({    'A': np.random.rand(num_samples),    'B': np.random.rand(num_samples),    'C': np.random.rand(num_samples),    'D': np.random.rand(num_samples),    'label': np.random.randint(0, 2, num_samples)})train, test = train_test_split(data, test_size=0.2, random_state=42) # 调整test_sizetrain_X = train[["A","B","C", "D"]].to_numpy()test_X = test[["A","B", "C", "D"]].to_numpy()train_Y = train[["label"]].to_numpy()test_Y = test[["label"]].to_numpy()train_X = torch.tensor(train_X, dtype=torch.float32)test_X = torch.tensor(test_X, dtype=torch.float32)train_Y = torch.tensor(train_Y, dtype=torch.float32) # 保持(N, 1)形状test_Y = torch.tensor(test_Y, dtype=torch.float32)   # 保持(N, 1)形状batch_size = 64train_dataset = TensorDataset(train_X, train_Y)# test_dataset = TensorDataset(test_X, test_Y) # 评估时通常直接使用test_X, test_Ytrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)# test_dataloader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False) # 如果需要批量评估,也可以使用# 2. 模型定义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)) # 输出范围0-1        return xinput_size = train_X.shape[1]hidden_size1 = 64hidden_size2 = 32output_size = 1 # 二分类输出model = SimpleClassifier(input_size,

以上就是PyTorch二分类模型准确率计算陷阱与修正:对比TensorFlow实践的详细内容,更多请关注创想鸟其它相关文章!

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

赞 (0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
python静态方法的用法
上一篇 2025年12月14日 15:27:43
优化 QLoRA 训练:解决大 Batch Size 导致训练时间过长的问题
下一篇 2025年12月14日 15:27:54

相关推荐

  • Java中ArrayList引用传递陷阱:避免数据意外修改的策略

    Java中ArrayList引用传递陷阱:避免数据意外修改的策略Java中ArrayList引用传递陷阱:避免数据意外修改的策略Java中ArrayList引用传递陷阱:避免数据意外修改的策略Java中ArrayList引用传递陷阱:避免数据意外修改的策略

    本文探讨了Java中ArrayList作为引用类型在对象构造时可能导致的数据意外修改问题。当将同一个ArrayList实例传递给多个对象后,对该列表的后续操作(如清空或添加元素)会影响所有引用它的对象。核心解决方案是为每个需要独立数据副本的对象,实例化一个新的ArrayList,从而确保数据隔离和一…

    2026年9月28日 • 用户投稿
    000
  • 豆包AI如何实现图像识别?教你搭建计算机视觉模型

    豆包AI如何实现图像识别?教你搭建计算机视觉模型豆包AI如何实现图像识别?教你搭建计算机视觉模型豆包AI如何实现图像识别?教你搭建计算机视觉模型豆包AI如何实现图像识别?教你搭建计算机视觉模型

    豆包ai本身不直接提供图像识别模型训练功能,但可结合第三方工具实现。1. 准备数据集:收集高质量、多样化的图像并划分训练集与验证集,或使用公开数据集。2. 搭建模型结构:采用迁移学习方法,选用resnet等预训练模型,调整输出层并加入防止过拟合的机制,豆包ai可生成代码框架。3. 训练与调参:设置合…

    2026年9月28日 • 用户投稿
    100
  • 武侠世界起航指南:从萌新到高手的全章节精要攻略

    武侠世界起航指南:从萌新到高手的全章节精要攻略武侠世界起航指南:从萌新到高手的全章节精要攻略武侠世界起航指南:从萌新到高手的全章节精要攻略武侠世界起航指南:从萌新到高手的全章节精要攻略

    踏入江湖的第一步,如何走稳走远?这份深度章节指南助你精准规划,避开弯路,高效解锁绝世武功与隐藏机缘! 第一章:初入江湖 – 筑基破局 核心目标: 击败管家 + 两名教头(新手战力检验) 与张风对话并切磋取胜(开启江湖路) 隐藏门派的钥匙(散人必看): 在朱宇处习得一气功(基础内功)!这是…

    2026年9月28日 • 用户投稿
    000
  • Android动态复选框状态持久化:SharedPreferences实践指南

    Android动态复选框状态持久化:SharedPreferences实践指南Android动态复选框状态持久化:SharedPreferences实践指南Android动态复选框状态持久化:SharedPreferences实践指南Android动态复选框状态持久化:SharedPreferences实践指南

    本教程详细阐述了如何在Android应用中持久化动态创建的复选框状态。通过利用SharedPreferences这一轻量级数据存储机制,我们能够确保用户在勾选或取消勾选动态生成的复选框后,其状态即使在应用重启或Activity重建后也能得以保留。文章将提供具体的代码示例和实现步骤,帮助开发者构建更具…

    2026年9月28日 • 用户投稿
    000
  • 笔尖AI语音识别不灵敏:灵敏度调整与方言适配技巧

    笔尖AI语音识别不灵敏:灵敏度调整与方言适配技巧笔尖AI语音识别不灵敏:灵敏度调整与方言适配技巧笔尖AI语音识别不灵敏:灵敏度调整与方言适配技巧笔尖AI语音识别不灵敏:灵敏度调整与方言适配技巧

    笔尖ai语音识别不灵敏可通过调整灵敏度、优化环境设置、进行方言适配等方式解决。首先,检查设置中的语音识别选项,通过滑块或数值逐步提高或降低灵敏度,根据使用场景选择合适的配置文件,并确保麦克风位置正确或更换高质量麦克风;其次,进行方言适配时,先检查语言设置是否有方言选项,若无则可自定义词汇并建立方言与…

    2026年9月28日 • 用户投稿
    000
  • Java集合引用管理:确保对象创建时内部列表状态独立的策略

    Java集合引用管理:确保对象创建时内部列表状态独立的策略Java集合引用管理:确保对象创建时内部列表状态独立的策略Java集合引用管理:确保对象创建时内部列表状态独立的策略Java集合引用管理:确保对象创建时内部列表状态独立的策略

    本教程探讨Java中将集合作为参数传递给构造函数时,如何避免因引用共享导致的内部数据意外更改问题。当多个对象共享同一个可变集合实例,并在外部修改该集合时,所有引用该集合的对象都会受影响。文章将详细介绍通过创建新集合实例或进行防御性复制两种有效策略,确保每个对象拥有独立且稳定的内部数据状态。 问题背景…

    2026年9月28日 • 用户投稿
    100
  • ChatSonic 创作 SEO 文案?关键词嵌入指令技巧​

    ChatSonic 创作 SEO 文案?关键词嵌入指令技巧​ChatSonic 创作 SEO 文案?关键词嵌入指令技巧​ChatSonic 创作 SEO 文案?关键词嵌入指令技巧​ChatSonic 创作 SEO 文案?关键词嵌入指令技巧​

    要写出高质量、能排名的 seo 文案,不能只依赖 chatsonic,还需掌握关键词嵌入技巧并对内容进行深度加工。1. 明确目标关键词与长尾关键词,专注几个核心词;2. 在 prompt 中明确指定关键词及出现位置,如标题、段首段尾等,但避免堆砌;3. 对生成内容进行润色,使其更自然流畅,并加入个人…

    2026年9月28日 • 用户投稿
    100
  • VSCode如何集成Git版本控制 VSCode中Git操作的便捷技巧

    首先确认git已安装并配置好用户名和邮箱;2. vscode通常自动检测git,若未检测到可手动在设置中指定git.path;3. 在vscode中打开项目并使用内置终端运行git init初始化仓库;4. 通过左侧源代码管理图标暂存、提交和推送更改;5. 遇到提交乱码时将files.encodin…

    2026年9月28日
    200
  • 和豆包一样的ai图片生成工具2025推荐top10

    2025年AI图片生成工具选择多样,boardmix因支持文生图、图生图、AI抠图、多种风格及在线协作,适合初学者与团队使用,且提供免费版,成为易用性高、功能全面的优选之一。 ☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模型☜☜☜ 2025年,想找个…

    2026年9月28日
    000
  • Android RecyclerView优化:通过DiffUtil实现增量更新

    Android RecyclerView优化:通过DiffUtil实现增量更新Android RecyclerView优化:通过DiffUtil实现增量更新Android RecyclerView优化:通过DiffUtil实现增量更新Android RecyclerView优化:通过DiffUtil实现增量更新

    本教程旨在解决RecyclerView在数据更新时(尤其是新增数据)出现的全量刷新和闪烁问题。通过详细介绍Android DiffUtil机制,我们将学习如何高效地进行列表项的增量更新,从而提升用户体验,避免不必要的UI重绘,特别适用于实时聊天等频繁数据变动的场景。 在开发Android应用时,Re…

    2026年9月28日 • 用户投稿
    100
  • sublime怎么配置clangd进行c++代码补全_Clangd插件C++环境配置

    sublime怎么配置clangd进行c++代码补全_Clangd插件C++环境配置sublime怎么配置clangd进行c++代码补全_Clangd插件C++环境配置sublime怎么配置clangd进行c++代码补全_Clangd插件C++环境配置sublime怎么配置clangd进行c++代码补全_Clangd插件C++环境配置

    配置Clangd实现C++智能补全,需安装LSP插件和Clangd服务器,并通过compile_commands.json告知编译信息,从而获得语义级代码补全、实时诊断与重构支持,显著提升Sublime Text的C++开发体验。 在Sublime Text里配置Clangd来搞定C++代码补全,说…

    2026年9月28日 • 用户投稿
    000
  • 豆包AI安装需要哪些运行时库 豆包AI系统依赖项完整清单

    豆包AI安装需要哪些运行时库 豆包AI系统依赖项完整清单豆包AI安装需要哪些运行时库 豆包AI系统依赖项完整清单豆包AI安装需要哪些运行时库 豆包AI系统依赖项完整清单豆包AI安装需要哪些运行时库 豆包AI系统依赖项完整清单

    #%#$#%@%@%$#%$#%#%#$%@_b05121b5eff2c++ee27d5b7d6a4dd8f2af运行需要python 3.8+、numpy、pandas、requests、torch/tensorflow、transformers、gradio/streamlit等核心库;操作系统…

    2026年9月28日 • 用户投稿
    100
  • 将Java或Groovy中的字符串转换为JSON对象

    将Java或Groovy中的字符串转换为JSON对象将Java或Groovy中的字符串转换为JSON对象将Java或Groovy中的字符串转换为JSON对象将Java或Groovy中的字符串转换为JSON对象

    将Java或Groovy中的字符串转换为JSON对象,需要根据实际情况进行分析。如果字符串是标准的JSON格式,可以直接使用JSON解析库进行转换。但如果字符串不是标准的JSON格式,则需要自定义解析器。 理解JSON格式 首先,我们需要明确标准的JSON格式。一个JSON对象是由键值对组成的,键和…

    2026年9月28日 • 用户投稿
    000
  • sublime代码提示不出来怎么办_解决Sublime代码自动补全失效问题

    sublime代码提示不出来怎么办_解决Sublime代码自动补全失效问题sublime代码提示不出来怎么办_解决Sublime代码自动补全失效问题sublime代码提示不出来怎么办_解决Sublime代码自动补全失效问题sublime代码提示不出来怎么办_解决Sublime代码自动补全失效问题

    代码提示失效多因插件未安装、语法识别错误或auto_complete被关闭。检查设置中是否启用auto_complete,安装Emmet、Anaconda等语言插件,确认文件语法正确,必要时清除缓存重建索引,可恢复补全功能。 Sublime Text 代码提示(自动补全)失效是不少用户在开发过程中遇…

    2026年9月28日 • 用户投稿
    400
  • 多模态AI可以生成视频吗 视频创作能力实测

    多模态AI可以生成视频吗 视频创作能力实测多模态AI可以生成视频吗 视频创作能力实测多模态AI可以生成视频吗 视频创作能力实测多模态AI可以生成视频吗 视频创作能力实测

    多模态ai确实能生成视频,但目前主要限于几秒到十几秒的短片段。其常见方式包括:1. 文本驱动生成,如输入描述生成森林日出画面;2. 图像扩展成视频,让静态图动态化;3. 图文混合引导生成更精准视频序列。当前生成视频存在长度有限、帧间不连贯、画质不稳定等问题,但适合社交媒体、创意样片等场景。建议创作者…

    2026年9月28日 • 用户投稿
    000
  • 数智融合驱动新质生产力,欧姆龙自动化亮相2025工博会

    数智融合驱动新质生产力,欧姆龙自动化亮相2025工博会数智融合驱动新质生产力,欧姆龙自动化亮相2025工博会数智融合驱动新质生产力,欧姆龙自动化亮相2025工博会数智融合驱动新质生产力,欧姆龙自动化亮相2025工博会

    作为全球自动化领域的数字化转型领军企业,欧姆龙自动化(中国)有限公司(以下简称“欧姆龙”)在第25届中国国际工业博览会精彩亮相。本次展会,欧姆龙精心打造了智能革新应用、数字驱动未来、强大产品矩阵三大主题展区,集中呈现多项契合现代制造业发展趋势的创新解决方案,为观众带来一场融合科技与智慧的智能制造盛宴…

    2026年9月28日 • 用户投稿
    500
  • 如何在Java中使用循环直到输入特定字符串?

    如何在Java中使用循环直到输入特定字符串?如何在Java中使用循环直到输入特定字符串?如何在Java中使用循环直到输入特定字符串?如何在Java中使用循环直到输入特定字符串?

    本文将解释如何在Java中使用while循环接收用户输入,并根据特定字符串(例如 “quit”)来终止循环。文章将解释为什么不能使用 == 运算符比较字符串,并提供使用 equals() 方法的正确示例,确保循环在用户输入特定字符串时正常退出。 在Java中,控制循环的执行直…

    2026年9月28日 • 用户投稿
    000
  • 如何在Jupyter中运行AI代码 Jupyter Notebook环境配置要点

    如何在Jupyter中运行AI代码 Jupyter Notebook环境配置要点如何在Jupyter中运行AI代码 Jupyter Notebook环境配置要点如何在Jupyter中运行AI代码 Jupyter Notebook环境配置要点如何在Jupyter中运行AI代码 Jupyter Notebook环境配置要点

    在jupyter notebook中运行ai代码的关键在于正确配置环境。1. 安装python 3.8+和pip,并通过命令行验证安装;2. 使用虚拟环境隔离项目依赖,激活后安装ai库如torch、tensorflow;3. 安装并启动jupyter notebook,必要时手动添加内核以确保其使用…

    2026年9月28日 • 用户投稿
    400
  • 谷歌浏览器如何给打开的标签页进行分组_谷歌浏览器标签页分组方法

    通过标签页分组功能可高效管理Chrome浏览器中大量标签,支持创建分组、添加标签页、自定义颜色名称、展开折叠及移除操作,提升浏览效率。 如果您在使用谷歌浏览器时打开了大量标签页,导致页面混乱难以管理,可以通过标签页分组功能将相关网页归类整理,提升浏览效率。以下是具体操作方法。 本文运行环境:MacB…

    2026年9月28日
    200
  • 前端验证后调用Servlet的正确方法

    前端验证后调用Servlet的正确方法前端验证后调用Servlet的正确方法前端验证后调用Servlet的正确方法前端验证后调用Servlet的正确方法

    本文旨在解决在前端JavaScript验证后如何正确调用Servlet的问题。通过分析常见的错误原因,例如表单提交事件的阻止和页面重载,以及Servlet中HTTP方法的使用,提供了一种清晰的解决方案,确保在前端验证通过后,能够成功地向Servlet发送请求并处理用户登录。 在Web开发中,经常需要…

    2026年9月28日 • 用户投稿
    300

发表回复

登录后才能评论
关注微信