turns-00020.parquet:76591
28e7e5e894ec4d03de195c4f
turn 2/16gpt-4-0125-previewChineseHong Kong1199 words
degenerate_repetitionAbsentFinal dense release
USER
当前代码① # 将 4输入分开,构建新的相同模态结合的2输入,2分支
import math
import logging
from functools import partial
from collections import OrderedDict
from copy import deepcopy
import torch
import torch.nn as nn
import torch.nn.functional as F
from timm.models.layers import to_2tuple
from lib.models.layers.patch_embed import PatchEmbed, PatchEmbed_event, xcorr_depthwise
from .utils import combine_tokens, recover_tokens
from .vit import VisionTransformer
from ..layers.attn_blocks import CEBlock
from .new_counter_guide import Counter_Guide
# from .ad_counter_guide import Counter_Guide_Enhanced
from .ad_counter_guide_downdim import Counter_Guide_Enhanced
_logger = logging.getLogger(__name__)
class VisionTransformerCE(VisionTransformer):
""" Vision Transformer with candidate elimination (CE) module
A PyTorch impl of : `An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale`
- https://arxiv.org/abs/2010.11929
Includes distillation token & head support for `DeiT: Data-efficient Image Transformers`
- https://arxiv.org/abs/2012.12877
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12,
num_heads=12, mlp_ratio=4., qkv_bias=True, representation_size=None, distilled=False,
drop_rate=0., attn_drop_rate=0., drop_path_rate=0., embed_layer=PatchEmbed, norm_layer=None,
act_layer=None, weight_init='',
ce_loc=None, ce_keep_ratio=None):
super().__init__()
if isinstance(img_size, tuple):
self.img_size = img_size
else:
self.img_size = to_2tuple(img_size)
self.patch_size = patch_size
self.in_chans = in_chans
self.num_classes = num_classes
self.num_features = self.embed_dim = embed_dim # num_features for consistency with other models
self.num_tokens = 2 if distilled else 1
norm_layer = norm_layer or partial(nn.LayerNorm, eps=1e-6)
act_layer = act_layer or nn.GELU
self.patch_embed = embed_layer(
img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim)
num_patches = self.patch_embed.num_patches
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.dist_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) if distilled else None
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + self.num_tokens, embed_dim))
self.pos_drop = nn.Dropout(p=drop_rate)
self.pos_embed_event = PatchEmbed_event(in_chans=32, embed_dim=768, kernel_size=4, stride=4)
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
blocks = []
ce_index = 0
self.ce_loc = ce_loc
for i in range(depth):
ce_keep_ratio_i = 1.0
if ce_loc is not None and i in ce_loc:
ce_keep_ratio_i = ce_keep_ratio[ce_index]
ce_index += 1
blocks.append(
CEBlock(
dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, drop=drop_rate,
attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer, act_layer=act_layer,
keep_ratio_search=ce_keep_ratio_i)
)
self.blocks = nn.Sequential(*blocks)
self.norm = norm_layer(embed_dim)
self.init_weights(weight_init)
# 添加交互模块counter_guide
self.counter_guide = Counter_Guide_Enhanced(768, 768)
def forward_features(self, z, x, event_z, event_x,
mask_z=None, mask_x=None,
ce_template_mask=None, ce_keep_rate=None,
return_last_attn=False
):
# 分支1 处理流程
B, H, W = x.shape[0], x.shape[2], x.shape[3]
x = self.patch_embed(x)
z = self.patch_embed(z)
z += self.pos_embed_z
x += self.pos_embed_x
if mask_z is not None and mask_x is not None:
mask_z = F.interpolate(mask_z[None].float(), scale_factor=1. / self.patch_size).to(torch.bool)[0]
mask_z = mask_z.flatten(1).unsqueeze(-1)
mask_x = F.interpolate(mask_x[None].float(), scale_factor=1. / self.patch_size).to(torch.bool)[0]
mask_x = mask_x.flatten(1).unsqueeze(-1)
mask_x = combine_tokens(mask_z, mask_x, mode=self.cat_mode)
mask_x = mask_x.squeeze(-1)
if self.add_cls_token:
cls_tokens = self.cls_token.expand(B, -1, -1)
cls_tokens = cls_tokens + self.cls_pos_embed
if self.add_sep_seg:
x += self.search_segment_pos_embed
z += self.template_segment_pos_embed
x = combine_tokens(z, x, mode=self.cat_mode)
if self.add_cls_token:
x = torch.cat([cls_tokens, x], dim=1)
x = self.pos_drop(x)
lens_z = self.pos_embed_z.shape[1]
lens_x = self.pos_embed_x.shape[1]
global_index_t = torch.linspace(0, lens_z - 1, lens_z).to(x.device)
global_index_t = global_index_t.repeat(B, 1)
global_index_s = torch.linspace(0, lens_x - 1, lens_x).to(x.device)
global_index_s = global_index_s.repeat(B, 1)
removed_indexes_s = []
# # 分支2 处理流程
event_x = self.pos_embed_event(event_x)
event_z = self.pos_embed_event(event_z)
event_x += self.pos_embed_x
event_z += self.pos_embed_z
event_x = combine_tokens(event_z, event_x, mode=self.cat_mode)
if self.add_cls_token:
event_x = torch.cat([cls_tokens, event_x], dim=1)
lens_z = self.pos_embed_z.shape[1]
lens_x = self.pos_embed_x.shape[1]
global_index_t1 = torch.linspace(0, lens_z - 1, lens_z).to(event_x.device)
global_index_t1 = global_index_t1.repeat(B, 1)
global_index_s1 = torch.linspace(0, lens_x - 1, lens_x).to(event_x.device)
global_index_s1 = global_index_s1.repeat(B, 1)
removed_indexes_s1 = []
for i, blk in enumerate(self.blocks):
# 第一个分支处理
x, global_index_t, global_index_s, removed_index_s, attn = \
blk(x, global_index_t, global_index_s, mask_x, ce_template_mask, ce_keep_rate)
# 第二个分支处理
event_x, global_index_t1, global_index_s1, removed_index_s1, attn1 = \
blk(event_x, global_index_t1, global_index_s1, mask_x, ce_template_mask, ce_keep_rate)
if self.ce_loc is not None and i in self.ce_loc:
removed_indexes_s.append(removed_index_s)
removed_indexes_s1.append(removed_index_s1)
# 在第1层增加counter_guide模块,验证早期融合效果
if i == 0 :
# 新增原始特征(在counter_guide之前的特征,便于后续的loss计算)
A_original_x = x.clone() # 保存原始x特征 (RGB)
A_original_event_x = event_x.clone() # 保存原始event特征 (event)
# 引入counter_guide进行特征增强
enhanced_x, enhanced_event_x = self.counter_guide(x, event_x)
# 保存经过counter_guide增强的特征
A_enhanced_x = enhanced_x # 保存增强后的特征
A_enhanced_event_x = enhanced_event_x # 保存增强后的特征
# 将增强后的特征与原特征相加
x = x + enhanced_x
event_x = event_x + enhanced_event_x
# 应用LayerNorm归一化处理
x = self.norm(x)
event_x = self.norm(event_x)
x_cat = torch.cat([event_x,x], dim=1)
x = x_cat
aux_dict = {
"attn": attn,
'attn1': attn1,
"removed_indexes_s": removed_indexes_s, # used for visualization
'removed_indexes_s1': removed_indexes_s1,
# 在aux_dict中返回原始特征和增强特征
'A_original_x': A_original_x,
'A_original_event_x': A_original_event_x,
'A_enhanced_x': A_enhanced_x,
'A_enhanced_event_x': A_enhanced_event_x,
}
return x, aux_dict
def forward(self, z, x, event_z, event_x,
ce_template_mask=None, ce_keep_rate=None,
tnc_keep_rate=None,
return_last_attn=False):
x, aux_dict = self.forward_features(z, x, event_z, event_x, ce_template_mask=ce_template_mask, ce_keep_rate=ce_keep_rate,)
return x, aux_dict
def _create_vision_transformer(pretrained=False, **kwargs):
model = VisionTransformerCE(**kwargs)
if pretrained:
if 'npz' in pretrained:
model.load_pretrained(pretrained, prefix='')
else:
checkpoint = torch.load(pretrained, map_location="cpu")
missing_keys, unexpected_keys = model.load_state_dict(checkpoint["model"], strict=False)
print('Load pretrained model from: ' + pretrained)
return model
def vit_base_patch16_224_ce(pretrained=False, **kwargs):
""" ViT-Base model (ViT-B/16) from original paper (https://arxiv.org/abs/2010.11929).
"""
model_kwargs = dict(
patch_size=16, embed_dim=768, depth=12, num_heads=12, **kwargs)
model = _create_vision_transformer(pretrained=pretrained, **model_kwargs)
return model
def vit_large_patch16_224_ce(pretrained=False, **kwargs):
""" ViT-Large model (ViT-L/16) from original paper (https://arxiv.org/abs/2010.11929).
"""
model_kwargs = dict(
patch_size=16, embed_dim=1024, depth=24, num_heads=16, **kwargs)
model = _create_vision_transformer(pretrained=pretrained, **model_kwargs)
return model和代码② from . import BaseActor
from lib.utils.misc import NestedTensor
from lib.utils.box_ops import box_cxcywh_to_xyxy, box_xywh_to_xyxy
import torch
from lib.utils.merge import merge_template_search
from ...utils.heapmap_utils import generate_heatmap
from ...utils.ce_utils import generate_mask_cond, adjust_keep_rate
#
import torch.nn.functional as F
class CEUTrackActor(BaseActor):
""" Actor for training CEUTrack models """
def __init__(self, net, objective, loss_weight, settings, cfg=None):
super().__init__(net, objective)
self.loss_weight = loss_weight
self.settings = settings
self.bs = self.settings.batchsize # batch size
self.cfg = cfg
def __call__(self, data):
"""
args:
data - The input data, should contain the fields 'template', 'search', 'gt_bbox'.
template_images: (N_t, batch, 3, H, W)
search_images: (N_s, batch, 3, H, W)
returns:
loss - the training loss
status - dict containing detailed losses
"""
# forward pass
out_dict = self.forward_pass(data)
# compute losses
loss, status = self.compute_losses(out_dict, data)
return loss, status
def forward_pass(self, data):
# currently only support 1 template and 1 search region
assert len(data['template_images']) == 1
assert len(data['search_images']) == 1
assert len(data['template_event']) == 1
assert len(data['search_event']) == 1
template_list = []
for i in range(self.settings.num_template):
template_img_i = data['template_images'][i].view(-1,
*data['template_images'].shape[2:]) # (batch, 3, 128, 128)
# template_att_i = data['template_att'][i].view(-1, *data['template_att'].shape[2:]) # (batch, 128, 128)
template_list.append(template_img_i)
search_img = data['search_images'][0].view(-1, *data['search_images'].shape[2:]) # (batch, 3, 320, 320)
# search_att = data['search_att'][0].view(-1, *data['search_att'].shape[2:]) # (batch, 320, 320)
template_event = data['template_event'][0].view(-1, *data['template_event'].shape[2:])
search_event = data['search_event'][0].view(-1, *data['search_event'].shape[2:])
box_mask_z = None
ce_keep_rate = None
if self.cfg.MODEL.BACKBONE.CE_LOC:
box_mask_z = generate_mask_cond(self.cfg, template_list[0].shape[0], template_list[0].device,
data['template_anno'][0])
ce_start_epoch = self.cfg.TRAIN.CE_START_EPOCH
ce_warm_epoch = self.cfg.TRAIN.CE_WARM_EPOCH
ce_keep_rate = adjust_keep_rate(data['epoch'], warmup_epochs=ce_start_epoch,
total_epochs=ce_start_epoch + ce_warm_epoch,
ITERS_PER_EPOCH=1,
base_keep_rate=self.cfg.MODEL.BACKBONE.CE_KEEP_RATIO[0])
if len(template_list) == 1:
template_list = template_list[0]
out_dict = self.net(template=template_list,
search=search_img,
event_template=template_event,
event_search=search_event,
ce_template_mask=box_mask_z,
ce_keep_rate=ce_keep_rate,
return_last_attn=False)
return out_dict
def compute_losses(self, pred_dict, gt_dict, return_status=True):
# gt gaussian map
gt_bbox = gt_dict['search_anno'][-1] # (Ns, batch, 4) (x1,y1,w,h) -> (batch, 4)
gt_gaussian_maps = generate_heatmap(gt_dict['search_anno'], self.cfg.DATA.SEARCH.SIZE, self.cfg.MODEL.BACKBONE.STRIDE)
gt_gaussian_maps = gt_gaussian_maps[-1].unsqueeze(1)
# Get boxes
pred_boxes = pred_dict['pred_boxes']
if torch.isnan(pred_boxes).any():
raise ValueError("Network outputs is NAN! Stop Training")
num_queries = pred_boxes.size(1)
pred_boxes_vec = box_cxcywh_to_xyxy(pred_boxes).view(-1, 4) # (B,N,4) --> (BN,4) (x1,y1,x2,y2)
gt_boxes_vec = box_xywh_to_xyxy(gt_bbox)[:, None, :].repeat((1, num_queries, 1)).view(-1, 4).clamp(min=0.0,
max=1.0) # (B,4) --> (B,1,4) --> (B,N,4)
# compute giou and iou
try:
giou_loss, iou = self.objective['giou'](pred_boxes_vec, gt_boxes_vec) # (BN,4) (BN,4)
except:
giou_loss, iou = torch.tensor(0.0).cuda(), torch.tensor(0.0).cuda()
# compute l1 loss
l1_loss = self.objective['l1'](pred_boxes_vec, gt_boxes_vec) # (BN,4) (BN,4)
# compute location loss
if 'score_map' in pred_dict:
location_loss = self.objective['focal'](pred_dict['score_map'], gt_gaussian_maps)
else:
location_loss = torch.tensor(0.0, device=l1_loss.device)
#######
# 计算特征一致性损失
feature_consistency_loss = F.mse_loss(pred_dict['A_enhanced_x'], pred_dict['A_original_x']) + F.mse_loss(pred_dict['A_enhanced_event_x'], pred_dict['A_original_event_x'])
# 标量化特征一致性损失
feature_consistency_loss = feature_consistency_loss.mean()
# weighted sum
loss = self.loss_weight['giou'] * giou_loss + self.loss_weight['l1'] * l1_loss + self.loss_weight['focal'] * location_loss + 0.5 * feature_consistency_loss
if return_status:
# status for log
mean_iou = iou.detach().mean()
status = {"Loss/total": loss.item(),
"Loss/giou": giou_loss.item(),
"Loss/l1": l1_loss.item(),
"Loss/location": location_loss.item(),
'Loss/feature_consistency': feature_consistency_loss.item(),
"IoU": mean_iou.item()}
return loss, status
else:
return loss
其中的feature_consistency_loss设计正确吗?ASSISTANT
代码②中的`feature_consistency_loss`计算方法基本是正确的,但是否完全适用取决于你的具体需求和目标。理论依据是距离度量(例如,均方误差MSE)通常用于评估两个特征之间的相似性或一致性。这里`feature_consistency_loss`是用来衡量经过特征增强(`counter_guide`模块处理后)的特征(`A_enhanced_x`, `A_enhanced_event_x`)与原始特征(`A_original_x`, `A_original_event_x`)之间的相似性。通过最小化这个损失,模型被激励在增强特征的同时,仍保留原始特征中重要的信息。 代码中的相应部分是: ```python feature_consistency_loss = F.mse_loss(pred_dict['A_enhanced_x'], pred_dict['A_original_x']) + F.mse_loss(pred_dict['A_enhanced_event_x'], pred_dict['A_original_event_x']) ``` 这段计算了RGB通道和事件通道经过`counter_guide`增强前后特征的MSE损失,并将这两个损失相加。这种做法希望模型在增强特征的同时不破坏原特征中的信息。 评价这种设计最重要的是考虑以下几点: 1. **目的明确**:这种特征一致性损失是否能够帮助你的模型更好地完成既定任务?在一些任务中,保持特征的一致性是很重要的,但在其他任务中,可能更需要的是提取更具有区分性的特征。 2. **权重选择**:最后,`feature_consistency_loss`加到总损失中的权重是`0.5`。这个权重是否合适,需要通过实验证明。模型的训练通常对损失函数中各项的权重比较敏感,这个权重可能需要通过实验来调整,以便找到最佳的设置。 3. **特性保留**:虽然这种做法可以帮助模型在增强特征的同时保留原有的特征信息,但也可能限制模型提取与任务密切相关但在原始数据中不明显的特征。是否使用这种方案需要根据实际情况和模型表现来判断。 4. **实验验证**:务必通过实验来验证这种设计的有效性,可以通过比较添加`feature_consistency_loss`前后模型的性能来做出判断。 总体来说,这种设计在概念上是合理的,但其效果和适用性需要通过具体的实验来确定,并根据实验结果进行调整。在实际应用中,还需要考虑不同损失项之间的平衡,确保模型能够兼顾增强特征的表达力和对原始特征信息的保留。