码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • mmsegmentation 添加L1Loss


    mmseg/models/losses/模块中添加L1Loss定义:

    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    from ..builder import LOSSES
    
    
    @LOSSES.register_module()
    class L1Loss(nn.Module):
        # TODO: weight
        def __init__(self, loss_name='loss_l1', **kwargs):
            super(L1Loss, self).__init__()
            self._loss_name = loss_name
    
        def forward(self, pred, target, weight=None, ignore_index=None):
       		# pred: (n,c,h,w)   target: (n,h,w)
            classes = pred.shape[1]
            size = list(target.shape)
            size.append(classes)  # (n,h,w,c)
            target_one_hot = target.view(-1)  # (n*h*w)
            ones = torch.sparse.torch.eye(classes).to(target_one_hot.device)
            ones = ones.index_select(0, target_one_hot)  # (n*h*w, classes)
            ones = ones.view(*size)  # (n,h,w,c)
            target_one_hot = ones.permute(0, 3, 1, 2)  # (n,c,h,w)
            loss = nn.L1Loss()(pred, target_one_hot)
            return loss
    
    	@property
        def loss_name(self):
            """Loss Name.
    
            This function must be implemented and will return the name of this
            loss function. This name will be used to combine different loss items
            by simple sum operation. In addition, if you want this loss item to be
            included into the backward graph, `loss_` must be the prefix of the
            name.
    
            Returns:
                str: The name of this loss item.
            """
            return self._loss_name
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21
    • 22
    • 23
    • 24
    • 25
    • 26
    • 27
    • 28
    • 29
    • 30
    • 31
    • 32
    • 33
    • 34
    • 35
    • 36
    • 37
    • 38
    • 39
    • 40

    注意必须要有loss_name方法,并且返回的loss_name需要以loss_作为前缀。

    传入的pred和target的shape不一致,需要转为一致才可以直接调用nn.L1Loss()方法。
    pred.shape: (n,c,h,w)
    target.shape: (n,h,w)
    所以需要将target转one-hot。转one-hot方法:index_select。
    (n,h,w) => (n,h,w,c) => (n,c,h,w)

    拓展阅读:Pytorch中,将label变成one hot编码的两种方式

  • 相关阅读:
    python+SQL sever+thinter学生宿舍管理系统
    【查找重复代码】python实现-附ChatGPT解析
    java基于Springboot+vue的宾馆酒店民宿入住管理平台系统 前后端分离elementui
    java计算机毕业设计高校在线办公系统源程序+mysql+系统+lw文档+远程调试
    BEVFormer代码跑通
    深入浅出PyTorch中的nn.CrossEntropyLoss
    RecyclerView Item中有EditText时点击事件处理
    人机逻辑中的家族相似性与非家族相似性
    C语言编写 输出[m,n]范围内所有“韩信点兵“数。
    Thread 和 Runnable 的区别
  • 原文地址:https://blog.csdn.net/qq_39735236/article/details/127806133
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号