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
# from .ad_counter_guide_multi import Counter_Guide_Enhanced_with_MultiHead
_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)
# self.counter_guide = Counter_Guide_Enhanced_with_MultiHead(768, 768, 3)
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 == 2 :
# 新增原始特征(在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([x,event_x], dim=1)
x_cat = event_x
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 ,然后加载的数据集:原始配置是:① import os
import os.path
import numpy as np
import torch
import csv
import pandas
import random
from collections import OrderedDict
from .base_video_dataset import BaseVideoDataset
from lib.train.data import jpeg4py_loader
from lib.train.admin import env_settings
import scipy.io as scio
class Coesot(BaseVideoDataset):
def __init__(self, root=None, image_loader=jpeg4py_loader, split=None, seq_ids=None, data_fraction=None):
root = env_settings().got10k_dir if root is None else root
super().__init__('Coesot', root, image_loader)
self.sequence_list = self._get_sequence_list()
# seq_id is the index of the folder inside the got10k root path
if split is not None:
if seq_ids is not None:
raise ValueError('Cannot set both split_name and seq_ids.')
if split == 'train':
file_path = os.path.join(self.root, 'train.txt')
elif split == 'val':
file_path = os.path.join(self.root, 'val.txt')
else:
raise ValueError('Unknown split name')
seq_ids = pandas.read_csv(file_path, header=None, dtype=np.int64).squeeze("columns").values.tolist()
elif seq_ids is None:
seq_ids = list(range(0, len(self.sequence_list)))
self.sequence_list = [self.sequence_list[i] for i in seq_ids]
def get_name(self):
return 'coesot'
def _get_sequence_list(self):
with open(os.path.join(self.root, 'list.txt')) as f:
dir_list = list(csv.reader(f))
dir_list = [dir_name[0] for dir_name in dir_list]
return dir_list
def _read_bb_anno(self, seq_path):
bb_anno_file = os.path.join(seq_path, "groundtruth.txt")
gt = pandas.read_csv(bb_anno_file, delimiter=',', header=None, dtype=np.float32, na_filter=False, low_memory=False).values
return torch.tensor(gt)
def _get_sequence_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id], self.sequence_list[seq_id] + "_aps")
def _get_event_img_sequence_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id], self.sequence_list[seq_id] + "_dvs")
def _get_grountgruth_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id])
def get_sequence_info(self, seq_id):
bbox_path = self._get_grountgruth_path(seq_id)
bbox = self._read_bb_anno(bbox_path)
valid = (bbox[:, 2] > 0) & (bbox[:, 3] > 0)
visible = valid.clone().byte()
# return {'bbox': bbox, 'valid': valid, 'visible': visible, 'visible_ratio': visible_ratio}
return {'bbox': bbox, 'valid': valid, 'visible': visible, }
def _get_frame_path(self, seq_path, frame_id):
if os.path.exists(os.path.join(seq_path, 'frame{:04}.png'.format(frame_id))):
return os.path.join(seq_path, 'frame{:04}.png'.format(frame_id)) # frames start from 0
else:
return os.path.join(seq_path, 'frame{:04}.bmp'.format(frame_id)) # some image is bmp
def _get_frame(self, seq_path, frame_id):
return self.image_loader(self._get_frame_path(seq_path, frame_id))
def _get_event_sequence_path(self, seq_id): ## get evemts' frames
return os.path.join(self.root, self.sequence_list[seq_id], self.sequence_list[seq_id] + "_voxel")
def _get_event_frame(self, seq_path, frame_id):
frame_event_list = []
for f_id in frame_id:
event_frame_file = os.path.join(seq_path, 'frame{:04}.mat'.format(f_id))
if os.path.getsize(event_frame_file) == 0:
event_features = np.zeros(4096, 19)
# need_data = [np.zeros([4096, 3]), np.zeros([4096, 16])]
else:
mat_data = scio.loadmat(event_frame_file)
# need_data = [mat_data['coor'], mat_data['features']]
event_features = np.concatenate((mat_data['coor'], mat_data['features']), axis=1) # concat coorelate and features (x,y,z, feauture32/16)
if np.isnan(event_features).any():
event_features = np.zeros(4096, 19)
print(event_frame_file, 'exist nan value in voxel.')
frame_event_list.append(event_features)
return frame_event_list
def get_frames(self, seq_id, frame_ids, anno=None):
seq_path = self._get_sequence_path(seq_id)
# obj_meta = self.sequence_meta_info[self.sequence_list[seq_id]]
frame_list = [self._get_frame(seq_path, f_id) for f_id in frame_ids]
seq_event_path = self._get_event_img_sequence_path(seq_id)
frame_event_img_list = [self._get_frame(seq_event_path, f_id) for f_id in frame_ids]
if anno is None:
anno = self.get_sequence_info(seq_id)
anno_frames = {}
for key, value in anno.items():
anno_frames[key] = [value[f_id, ...].clone() for f_id in frame_ids]
object_meta = OrderedDict({'object_class_name': None,
'motion_class': None,
'major_class': None,
'root_class': None,
'motion_adverb': None})
seq_event_path = self._get_event_sequence_path(seq_id)
frame_event_list = self._get_event_frame(seq_event_path, frame_ids)
return frame_list, anno_frames, object_meta, frame_event_list, frame_event_img_list
我们修改为:② import os.path
import numpy as np
import torch
import csv
import pandas
import random
from collections import OrderedDict
from .base_video_dataset import BaseVideoDataset
from lib.train.data import jpeg4py_loader
from lib.train.admin import env_settings
import scipy.io as scio
class Coesot(BaseVideoDataset):
def __init__(self, root=None, image_loader=jpeg4py_loader, split=None, seq_ids=None, data_fraction=None):
root = env_settings().coesot_dir if root is None else root
super().__init__('Coesot', root, image_loader)
self.sequence_list = self._get_sequence_list()
if split is not None:
if seq_ids is not None:
raise ValueError('Cannot set both split_name and seq_ids.')
if split == 'train':
file_path = os.path.join(self.root, 'train.txt')
elif split == 'val':
file_path = os.path.join(self.root, 'val.txt')
else:
raise ValueError('Unknown split name')
seq_ids = pandas.read_csv(file_path, header=None, dtype=np.int64).squeeze('columns').values.tolist()
elif seq_ids is None:
seq_ids = list(range(0, len(self.sequence_list)))
self.sequence_list = [self.sequence_list[i] for i in seq_ids]
def get_name(self):
return 'coesot'
def _get_sequence_list(self):
with open(os.path.join(self.root, 'list.txt')) as f:
dir_list = list(csv.reader(f))
dir_list = [dir_name[0] for dir_name in dir_list]
return dir_list
def _read_bb_anno(self, seq_path):
bb_anno_file = os.path.join(seq_path, 'groundtruth.txt')
gt = pandas.read_csv(bb_anno_file, delimiter=',', header=None, dtype=np.float32, na_filter=False, low_memory=False).values
return torch.tensor(gt)
def _get_sequence_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id], self.sequence_list[seq_id] + '_aps')
def _get_event_img_sequence_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id], self.sequence_list[seq_id] + '_dvs')
def _get_grountgruth_path(self, seq_id):
return os.path.join(self.root, self.sequence_list[seq_id])
def get_sequence_info(self, seq_id):
bbox_path = self._get_grountgruth_path(seq_id)
bbox = self._read_bb_anno(bbox_path)
valid = (bbox[:, 2] > 0) & (bbox[:, 3] > 0)
visible = valid.clone().byte()
return {'bbox': bbox, 'valid': valid, 'visible': visible}
def _get_frame_path(self, seq_path, frame_id):
if os.path.exists(os.path.join(seq_path, 'frame{:04}.png'.format(frame_id))):
return os.path.join(seq_path, 'frame{:04}.png'.format(frame_id)) # frames start from 0
else:
return os.path.join(seq_path, 'frame{:04}.bmp'.format(frame_id)) # some image is bmp
def _get_frame(self, seq_path, frame_id,anno=None):
return self.image_loader(self._get_frame_path(seq_path, frame_id))
def get_frames(self, seq_id, frame_ids, anno=None):
seq_path = self._get_sequence_path(seq_id)
seq_event_path = self._get_event_img_sequence_path(seq_id)
frame_list = [self._get_frame(seq_path, f_id) for f_id in frame_ids]
frame_event_img_list = [self._get_frame(seq_event_path, f_id) for f_id in frame_ids]
if anno is None:
anno = self.get_sequence_info(seq_id)
anno_frames = {}
for key, value in anno.items():
anno_frames[key] = [value[f_id, ...].clone() for f_id in frame_ids]
object_meta = OrderedDict({'object_class_name': None,
'motion_class': None,
'major_class': None,
'root_class': None,
'motion_adverb': None})
return frame_list, frame_event_img_list, anno_frames, object_meta ,有什么不同和影响吗?