半监督语义分割学习笔记

目录

[partial cross entropy loss](#partial cross entropy loss)


GitHub - LiheYoung/UniMatch: [CVPR 2023] Revisiting Weak-to-Strong Consistency in Semi-Supervised Semantic Segmentation

partial cross entropy loss

python 复制代码
import torch
import torch.nn.functional as F

def partial_cross_entropy_loss(inputs, targets, ignore_index=-1):
    """
    自定义部分交叉熵损失函数,忽略 ignore_index 指定的标签。
    
    :param inputs: 模型的输出,形状应为 (N, C, H, W),其中 N 是批量大小,C 是类别数,H 和 W 是高度和宽度。
    :param targets: 真实的标签,形状应为 (N, H, W)。
    :param ignore_index: 要忽略的标签值,默认为 -1。
    :return: 计算得到的损失。
    """
    # 计算 log softmax
    log_probs = F.log_softmax(inputs, dim=1)
    
    # 将 log_probs 和 targets 转换为适合 gather 的形状
    log_probs = log_probs.permute(0, 2, 3, 1)  # (N, H, W, C)
    log_probs = log_probs.reshape(-1, log_probs.shape[-1])  # (N*H*W, C)
    targets = targets.view(-1)  # (N*H*W)
    
    # 掩码未标记的数据点
    mask = targets != ignore_index
    log_probs = log_probs[mask]
    targets = targets[mask]
    
    # 只计算有标签的数据点的损失
    loss = F.nll_loss(log_probs, targets, reduction='mean')
    
    return loss
python 复制代码
# 假设模型的输出和真实标签
outputs = torch.randn(2, 3, 5, 5)  # 随机生成模拟输出(2个样本,3个类别,5x5的图像)
targets = torch.tensor([[[-1, 1, -1, 0, -1], 
                         [1, -1, 2, 2, 1], 
                         [-1, -1, 1, -1, 0], 
                         [2, 2, 2, -1, 1], 
                         [-1, 0, -1, 0, 1]], 
                        [[1, 0, -1, 1, -1], 
                         [2, 2, -1, 0, 0], 
                         [-1, 1, 1, 0, -1], 
                         [0, 0, 2, -1, 1], 
                         [2, -1, 0, -1, -1]]])  # 生成带有未标记区域的标签

# 计算损失
loss = partial_cross_entropy_loss(outputs, targets)
print(f"Loss: {loss.item()}")
相关推荐
ClutchoQ19 分钟前
【你指的API是哪个API?软件工程师跨服聊天实录】
笔记·其他
二哈赛车手2 小时前
新人笔记---Spring AI的Advisor以及其底层机制讲解(涉及源码),包含一些遇见的Spring AI的Advisor缺陷问题的解决方案
java·人工智能·spring boot·笔记·spring
red_redemption3 小时前
自由学习记录(181)
学习
wuxinyan1233 小时前
大模型学习之路007:RAG 零基础入门教程(第四篇):生成侧核心技术与大模型集成
人工智能·学习·rag
阿豪只会阿巴4 小时前
【没事学点啥】TurboBlog轻量级个人博客项目——Turbo Blog 项目学习与上线指南
开发语言·python·学习·状态模式
Slow菜鸟4 小时前
Docker 学习篇(三)| Docker安装指南(Linux版)
linux·学习·docker
Tutankaaa4 小时前
知识竞赛软件SaaS版 vs 本地部署
人工智能·经验分享·笔记·学习
小仙女的小稀罕4 小时前
培训要点写不完不会整理?规范培训转待办可这样操作
大数据·人工智能·学习·自然语言处理·语音识别
许长安4 小时前
RPC 异步调用基本使用方法:基于官方helloworld-async 示例
c++·经验分享·笔记·rpc
Wallace Zhang5 小时前
SimpleFOC源码学习10(v2.3.2) - 电流传感器CurrentSense.cpp与CurrentSense.h
驱动开发·stm32·学习·电流环·simplefoc·foc电机控制