AlibabaResearch / AlibabaResearch/efficientteacher

修改yolov8发现缺少loss文件

Abierto
#125 30 comentarios 0 reacciones 0 asignados Ver en GitHub
Lenguaje dominante
Python
Estrellas
813
Forks
126
Métricas de merge de PR
Sin PR fusionados en 30 d

Descripción

下面是tal_loss.py文件的显示, YOLOATSSAssigner,from models.loss.gfocal_loss import VarifocalLoss, BboxLoss这三个文件请问在哪里呢?我找了一圈没找到在哪?

`#!/usr/bin/env python3
# -*- coding:utf-8 -*-

import torch
import torch.nn as nn
import numpy as np
import torch.nn.functional as F
from models.module.nanodet_utils import generate_anchors
from models.module.nanodet_utils import dist2bbox, xywh2xyxy
# from loss.yolox_loss import IOUloss
from models.assigner.yolo_atss_assigner import YOLOATSSAssigner
from models.assigner.tal_assigner import TaskAlignedAssigner
from utils.torch_utils import is_parallel
from models.loss.gfocal_loss import VarifocalLoss, BboxLoss

class ComputeTalLoss:
'''Loss computation func.'''
def __init__(self,
model,
cfg):
# fpn_strides=[8, 16, 32]
# grid_cell_size=5.0
# grid_cell_offset=0.5
# num_classes=80,
# ori_img_size=640
# use_dfl=True
# reg_max=16
# iou_type='siou'
# num_classes = cfg.Dataset.nc
# loss_weight={ 'class': 1.0, 'iou': 2.5, 'dfl': 0.5}
# device = next(model.parameters()).device # get model device
self.epoch = 0
self.det = model.module.head if is_parallel(model) else model.head# Detect() module
self.fpn_strides = cfg.Model.Head.strides
self.grid_cell_size = cfg.Loss.grid_cell_size
self.grid_cell_offset = cfg.Loss.grid_cell_offset
self.num_classes = cfg.Dataset.nc
self.ori_img_size = cfg.Dataset.img_size

# warmup_epoch=4
self.warmup_epoch = cfg.hyp.warmup_epochs
self.warmup_assigner = YOLOATSSAssigner(9, num_classes=self.num_classes)
self.formal_assigner = TaskAlignedAssigner(top_k=13, num_classes=self.num_classes, alpha=1.0, beta=6.0)

self.use_dfl = cfg.Loss.use_dfl
self.use_gfl = cfg.Loss.use_gfl
self.reg_max = cfg.Loss.reg_max
self.iou_type = cfg.Loss.iou_type
self.proj = nn.Parameter(torch.linspace(0, self.reg_max, self.reg_max + 1), requires_grad=False)
self.varifocal_loss = VarifocalLoss().cuda()
self.bce = nn.BCELoss(reduction='none')
self.bbox_loss = BboxLoss(self.num_classes, self.reg_max, self.use_dfl, self.iou_type).cuda()
self.loss_weight = {'class': cfg.Loss.qfl_loss_weight, 'iou': cfg.Loss.box_loss_weight, 'dfl':cfg.Loss.dfl_loss_weight} `

Guía de contribución

No hay ninguna guía de contribución indexada para este repositorio

Evaluación

Este issue todavía no se ha evaluado.

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.