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
处理不同形状批次的损失计算:加权平均方法_创想鸟

处理不同形状批次的损失计算:加权平均方法

处理不同形状批次的损失计算:加权平均方法

引言

正如摘要所述,当处理形状不规则的批次数据时,损失计算需要特别处理。简单地平均每个样本的损失可能会导致偏差,因为较小的批次会与较大的批次产生相同的影响。为了解决这个问题,我们可以使用加权平均,根据每个批次的大小来调整其对整体损失的贡献。

问题描述

在训练过程中,如果每个批次的样本具有不同的长度或形状,则直接堆叠每个样本的损失并计算平均值可能会导致问题。例如,在序列数据处理中,每个序列的长度可能不同,因此每个批次中有效数据的数量也不同。以下代码展示了这个问题:

def training():    model.train()    train_mae = []    progress = tqdm(train_dataloader, desc='Training')    for batch_index, batch in enumerate(progress):        x = batch['x'].to(device)        x_lengths = batch['x_lengths'].to(device)        y = batch['y'].to(device)        y_type = batch['y_type'].to(device)        y_valid_indices = batch['y_valid_indices'].to(device)        # Zero Gradients        optimizer.zero_grad()        # Forward pass        y_first, y_second = model(x)        losses = []        for j in range(len(x_lengths)):            x_length = x_lengths[j].item()            if y_type[j].item() == 0:                predicted = y_first[j]            else:                predicted = y_second[j]            actual = y[j]            valid_mask = torch.zeros_like(predicted, dtype=torch.bool)            valid_mask[:x_length] = 1            # Padding of -1 is removed from y            indices_mask = y[j].ne(-1)            valid_indices = y[j][indices_mask]            valid_predicted = predicted[valid_mask]            valid_actual = actual[valid_mask]            loss = mae_fn(valid_predicted, valid_actual, valid_indices)            losses.append(loss)        # Backward pass and update        loss = torch.stack(losses).mean()   # This fails due to different shapes        loss.backward()        optimizer.step()        train_mae.append(loss.detach().cpu().numpy())        progress.set_description(            f"mae: {loss.detach().cpu().numpy():.4f}"        )    # Return the average MAEs for y type    return (        np.mean(train_mae)    )

在上述代码中,loss = torch.stack(losses).mean() 这一行会因为 losses 列表中的张量形状不同而失败。

解决方案:加权平均

为了解决这个问题,我们可以计算每个批次的平均损失,然后根据批次大小对这些平均损失进行加权平均。这样,较大的批次将对最终损失产生更大的影响,从而更准确地反映模型的性能。

以下是一个示例代码:

import torch# 示例数据losses_perbatch = [torch.randn(8, 1), torch.randn(4, 1), torch.randn(2, 1)]# 加权平均total_samples = sum([len(batch) for batch in losses_perbatch])weighted_mean_perbatch = torch.tensor([batch.sum() for batch in losses_perbatch]) / total_samples# 或者等价于:# weighted_mean_perbatch = torch.tensor([batch.mean() * len(batch) for batch in losses_perbatch]) / total_samplesfinal_weighted_loss = sum(weighted_mean_perbatch)print(f"Final Weighted Loss: {final_weighted_loss}")

在这个例子中,losses_perbatch 包含不同大小的批次的损失。我们首先计算所有批次的总样本数 total_samples。然后,对于每个批次,我们计算其损失的总和,并将其除以 total_samples,得到加权平均损失。最后,我们将所有批次的加权平均损失相加,得到最终的加权损失。

代码集成

将加权平均方法集成到原始的训练函数中,可以修改如下:

def training():    model.train()    train_mae = []    progress = tqdm(train_dataloader, desc='Training')    for batch_index, batch in enumerate(progress):        x = batch['x'].to(device)        x_lengths = batch['x_lengths'].to(device)        y = batch['y'].to(device)        y_type = batch['y_type'].to(device)        y_valid_indices = batch['y_valid_indices'].to(device)        # Zero Gradients        optimizer.zero_grad()        # Forward pass        y_first, y_second = model(x)        losses = []        batch_sizes = []  # Store the size of each batch        for j in range(len(x_lengths)):            x_length = x_lengths[j].item()            if y_type[j].item() == 0:                predicted = y_first[j]            else:                predicted = y_second[j]            actual = y[j]            valid_mask = torch.zeros_like(predicted, dtype=torch.bool)            valid_mask[:x_length] = 1            # Padding of -1 is removed from y            indices_mask = y[j].ne(-1)            valid_indices = y[j][indices_mask]            valid_predicted = predicted[valid_mask]            valid_actual = actual[valid_mask]            loss = mae_fn(valid_predicted, valid_actual, valid_indices)            losses.append(loss)            batch_sizes.append(x_length)  # Store the batch size        # Calculate weighted loss        total_samples = sum(batch_sizes)        weighted_mean_perbatch = torch.tensor([loss.sum() for loss in losses]) / total_samples        loss = sum(weighted_mean_perbatch)        # Backward pass and update        loss.backward()        optimizer.step()        train_mae.append(loss.detach().cpu().numpy())        progress.set_description(            f"mae: {loss.detach().cpu().numpy():.4f}"        )    # Return the average MAEs for y type    return (        np.mean(train_mae)    )

在这个修改后的代码中,我们添加了一个 batch_sizes 列表来存储每个批次的大小。然后,我们使用这些大小来计算加权平均损失,并将其用于反向传播和优化。

注意事项

确保 batch_sizes 列表中的大小与 losses 列表中的损失对应。加权平均方法可以更稳定地计算损失,但可能需要更多的计算资源。这种方法特别适用于处理序列数据或其他具有不同形状的批次数据。

总结

当处理不同形状的批次数据时,加权平均是一种有效的损失计算方法。通过考虑每个批次的大小,我们可以更准确地评估模型的性能,并避免简单平均可能导致的偏差。这种方法可以应用于各种机器学习任务,特别是那些涉及序列数据或其他形状不规则的数据的任务。

以上就是处理不同形状批次的损失计算:加权平均方法的详细内容,更多请关注创想鸟其它相关文章!

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

赞 (0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
处理不同形状批次的损失计算:加权平均损失方法
上一篇 2025年12月14日 10:43:18
Python中正确处理数据库NULL值:类型判断与转换
下一篇 2025年12月14日 10:43:30

相关推荐

  • 如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程

    如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程如何用PhotoLab的AI裁剪图片?快速实现智能图像裁剪教程

    PhotoLab的AI裁剪功能通过智能识别主体与构图原则,提供优化裁剪建议,区别于传统手动裁剪的纯物理操作,能自动应用美学法则提升照片视觉吸引力;在人像、社交媒体适配、风景静物等场景中表现突出,尤其擅长保留核心焦点并适配多平台比例;用户可导入图片后使用AI裁剪工具,系统分析画面并生成建议裁剪框,支持…

    2026年9月22日 • 用户投稿
    000
  • 递归实现列表排序检查与条件移除最大值

    本文详细介绍了如何使用Java递归方法处理整数列表。核心内容包括:首先检查列表是否已排序,如果已排序则直接返回false;如果未排序,则查找列表中的最大值。仅当最大值位于列表的起始或结束位置时,才将其移除并递归地继续处理列表。如果最大值位于列表中间,则打印当前列表并终止递归。 在数据处理和算法设计中…

    2026年9月22日
    000
  • 国泰航空“广州始发礼遇”限时开启,新增广州往返香港航班助力畅游亚洲

    落地即启程,中转再提速,轻松畅游亚洲 秋意正浓,正是踏上旅途的好时机。国泰航空为大湾区“1小时生活圈”注入全新活力——自2025年10月27日起,广州与中国香港之间的往返航班将加密至每日三班,并同步推出“广州始发礼遇”限时优惠活动,让旅客以更实惠的价格畅行亚洲热门目的地。 限时优惠抢先订 亚洲美景随…

    2026年9月22日
    100
  • VSCode如何实现代码可视化调试 VSCode执行流程图形化分析方法

    vscode的可视化调试功能通过内置调试器和扩展生态,显著提升代码理解与问题排查效率。1. 首先配置launch.json文件以定义调试环境,支持多种语言如node.js、python等;2. 在代码中设置断点,程序运行至断点时暂停,便于检查变量状态和执行上下文;3. 利用调试面板查看变量、监视表达…

    2026年9月22日
    000
  • MySQL备份压缩与加密技巧_MySQL提升备份安全与效率

    MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率MySQL备份压缩与加密技巧_MySQL提升备份安全与效率

    mysql备份压缩与加密的核心在于减少存储空间并提升数据安全性。1. 压缩能显著降低存储成本,提升传输效率,加快恢复速度,简化备份管理,并有助于满足合规要求;2. 加密则通过防止未授权访问保障数据安全。实现方式主要有:1. 使用mysqldump结合gzip和gpg/openssl进行逻辑备份、压缩…

    2026年9月22日 • 用户投稿
    100
  • VS Code中Dockerized PHP项目:解决PHP版本冲突的教程

    本教程旨在解决在VS Code中开发Dockerized PHP项目时,VS Code默认识别宿主机PHP版本而非容器内PHP版本的问题。核心解决方案是利用VS Code的Remote – Containers扩展,实现直接在Docker容器内部进行代码开发,从而确保VS Code及其所…

    2026年9月22日
    200
  • 蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!

    蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!蔡司2亿影像大小王,年度影像旗舰vivo X300系列发布!

    PConline最新资讯,vivo于今晚正式揭晓X300系列新机,定位“全焦段影像旗舰”,起售价为4399元。该系列成为首款搭载联发科天玑9500芯片的智能手机,并携手三星与索尼共同定制多颗影像传感器,在影像能力、屏幕素质及续航表现上力求全面跃升。 产品线涵盖X300与X300 Pro两款机型,价格…

    2026年9月22日 • 用户投稿
    000
  • 从AI场景搭建到蝴蝶号运营,全流程实战攻略

    从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略从AI场景搭建到蝴蝶号运营,全流程实战攻略

    做ai内容变现需先明确方向再选工具,注册蝴蝶号要模拟真实行为,用ai提升效率但需调整内容细节,流量转化重于播放量。一、先确定内容类型和风格,根据方向选择合适ai工具链搭建流程,用免费api测试效果。二、蝴蝶号注册尽量用企业主体,资料完整,养号阶段关注同类账号,保持每天发布1~2条内容,视频控制在30…

    2026年9月22日 • 用户投稿
    100
  • 优化Spring Boot应用:构建高效通用的DTO与实体映射服务

    本文旨在解决Spring Boot项目中DTO与实体间重复映射的痛点。通过引入一个基于泛型的抽象服务层,结合ModelMapper工具,我们展示了如何构建一个类型安全、可重用的通用映射机制。此方案显著减少了样板代码,提升了代码的可维护性和开发效率,避免了手动类型转换的繁琐与潜在错误。 在构建基于sp…

    2026年9月22日
    100
  • GIMP中如何利用AI裁剪图片?一步步完成高效图像裁剪方法

    GIMP虽无“一键AI裁剪”功能,但可通过智能选择工具(如前景选择、智能剪刀)精准选中主体,结合Resynthesizer插件的内容感知填充实现类AI裁剪效果;对于更高要求,可协同Remove.bg等外部AI工具完成自动抠图,再导入GIMP进行裁剪或背景替换,形成高效智能裁剪工作流。 ☞☞☞AI 智…

    2026年9月22日
    100
  • QQ音乐自动续费怎么停止_QQ音乐停止自动续费的详细步骤

    首先需手动取消自动续费,1.在QQ音乐App“我的”-“会员中心”-“个人中心”-“管理自动续费”中关闭;2.通过微信“服务”-“钱包”-“支付设置”-“自动续费”关闭QQ音乐会员;3.iOS用户需在“设置”-Apple ID-“订阅”中取消QQ音乐订阅,确认后当前周期结束即停止扣费。 如果您在使用…

    2026年9月22日
    200
  • 疑似荣耀500系列入网 代号Merry全系支持80W有线快充

    10月25日,知名数码博主“数码闲聊站”透露,荣耀500系列新机已现身工信部,型号分别为mep-an00和mey-an00,预计代号为merry/merryp,全系支持80w有线快充。该博主还表示,此前上手的样机提供了黑色、银色、粉色和蓝色等多种配色方案,外观设计或将延续前代爆款风格。 据最新消息,…

    2026年9月22日
    000
  • PHP each() 函数的替代方案:自定义实现与常见错误修正

    本文探讨了PHP中已废弃的each()函数的替代方案。针对常见的自定义实现,如myEach(),文章详细指出了其在返回数组结构中常犯的错误,并提供了正确的代码示例,以确保替代函数能够模拟each()的预期行为,帮助开发者编写更健壮、兼容未来的PHP代码。 理解 each() 函数及其废弃背景 在PH…

    2026年9月22日
    000
  • MAC怎么在登录界面显示自定义信息_macOS锁屏界面显示个性化文本

    1、通过系统设置可直接在登录界面显示自定义文本,进入“隐私与安全性”→“登录窗口”编辑消息;2、使用终端命令sudo defaults write写入LoginWindowText实现相同效果;3、企业可通过.mobileconfig描述文件集中部署登录信息。 如果您希望在Mac的登录界面显示个性化…

    2026年9月22日
    000
  • Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析

    Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析Vision Transformer 必读系列之图像分类综述(三): MLP、ConvMixer 和架构分析

    号外号外!awesome-vit 上新啦, 欢迎大家 Star Star Star ~ https://github.com/open-mmlab/awesome-vit 前言 在 Vision Transformer 必读系列之图像分类综述(一):概述 一文中对 Vision Transforme…

    2026年9月22日 • 用户投稿
    200
  • 蝴蝶号无人直播完整流程详解:搭建+开播+引流

    蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流蝴蝶号无人直播完整流程详解:搭建+开播+引流

    蝴蝶号无人直播的完整流程包括前期准备、直播搭建、开播设置、引流推广、监控与维护五个步骤。前期准备需完成账号注册认证、硬件设备配置、软件安装及素材准备;直播搭建涉及场景设置、素材导入、循环播放设定及自动化脚本配置;开播设置包括直播间信息填写、推流配置与测试直播;引流推广可通过平台内工具、社交媒体、内容…

    2026年9月22日 • 用户投稿
    100
  • 那些为小米信仰充值的人 都怎么样了?

    那些为小米信仰充值的人 都怎么样了?那些为小米信仰充值的人 都怎么样了?那些为小米信仰充值的人 都怎么样了?那些为小米信仰充值的人 都怎么样了?

    2024年末,小米的股价一路上扬,逼近40港元。而在此前的很长一段时间,外界因对小米造车的不信任,唱空小米,股价一度跌至10港元以下。 为了庆祝小米重回股价峰值,在一个寒气逼人的冬日,一群小米股民相聚在北京小米互联网园区的门口。 他们像个孩子一样,打出一条“心里有火,眼里有光”的横幅。一位教授喝到尽…

    2026年9月22日 • 用户投稿
    100
  • 如何在VEED.io中制作AI视频?在线工具快速剪辑AI内容的步骤

    如何在VEED.io中制作AI视频?在线工具快速剪辑AI内容的步骤如何在VEED.io中制作AI视频?在线工具快速剪辑AI内容的步骤如何在VEED.io中制作AI视频?在线工具快速剪辑AI内容的步骤如何在VEED.io中制作AI视频?在线工具快速剪辑AI内容的步骤

    VEED.io通过“文本转视频”和“AI形象”功能,让视频制作变得简单高效。用户只需输入文本,即可生成带AI配音、字幕和匹配素材的视频,或选择AI虚拟人物进行口型同步播报。平台还提供AI语音合成、自动字幕、多语言支持及丰富编辑功能,便于后期精修。优化效果需从高质量文本入手,合理选择声音与形象,并通过…

    2026年9月22日 • 用户投稿
    000
  • Java中递归处理列表:条件性移除最大值策略与实现

    本教程深入探讨了如何在Java中使用递归方法,根据特定条件(如列表是否已排序、最大值是否位于列表的首尾)来移除列表中的最大值。文章将详细阐述如何设计一个高效的递归算法,包括排序检查、最大值定位以及条件性移除的实现细节,并提供完整的代码示例和注意事项,帮助读者掌握递归在复杂列表操作中的应用。 引言:递…

    2026年9月22日
    000
  • 玩转 Spring Boot 集成篇(定时任务框架Quartz)

    玩转 Spring Boot 集成篇(定时任务框架Quartz)玩转 Spring Boot 集成篇(定时任务框架Quartz)玩转 Spring Boot 集成篇(定时任务框架Quartz)玩转 Spring Boot 集成篇(定时任务框架Quartz)

    在日常项目研发中,定时任务可谓是必不可少的一环,关于 spring boot 如何实现静态定时任务、动态定时任务以及如何开启多线程跑任务,均已在上篇分享过,不再赘述。 虽然 Spring Boot 内置注解方式实现的定时任务,在一定程度上也能解决一定的业务场景问题,但是若做更复杂的动作,例如启停任务…

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

发表回复

登录后才能评论
关注微信