| import torch |
| import torch.nn as nn |
| from mmcv.cnn import Linear, bias_init_with_prob, Scale |
|
|
| from mmcv.runner import force_fp32 |
| from mmdet.core import (build_assigner, build_sampler, multi_apply, |
| reduce_mean) |
| from mmdet.models.utils import build_transformer |
| from mmdet.models import HEADS, build_loss |
| from mmdet.models.dense_heads.anchor_free_head import AnchorFreeHead |
| from mmdet.models.utils.transformer import inverse_sigmoid |
| from mmdet3d.core.bbox.coders import build_bbox_coder |
| from projects.mmdet3d_plugin.core.bbox.util import normalize_bbox |
| from projects.mmdet3d_plugin.models.utils.positional_encoding import pos2posemb3d, pos2posemb1d, nerf_positional_encoding |
| from projects.mmdet3d_plugin.models.utils.misc import MLN, topk_gather, transform_reference_points, memory_refresh, SELayer_Linear |
| from projects.mmdet3d_plugin.models.routing import ( |
| GumbelTokenRouter, |
| apply_dense_budget_mask, |
| make_budget_mask, |
| ) |
| import copy |
| from mmdet.models.utils import NormedLinear |
|
|
| @HEADS.register_module() |
| class SparseHead(AnchorFreeHead): |
| """Implements the DETR transformer head. |
| See `paper: End-to-End Object Detection with Transformers |
| <https://arxiv.org/pdf/2005.12872>`_ for details. |
| Args: |
| num_classes (int): Number of categories excluding the background. |
| in_channels (int): Number of channels in the input feature map. |
| num_query (int): Number of query in Transformer. |
| num_reg_fcs (int, optional): Number of fully-connected layers used in |
| `FFN`, which is then used for the regression head. Default 2. |
| transformer (obj:`mmcv.ConfigDict`|dict): Config for transformer. |
| Default: None. |
| sync_cls_avg_factor (bool): Whether to sync the avg_factor of |
| all ranks. Default to False. |
| positional_encoding (obj:`mmcv.ConfigDict`|dict): |
| Config for position encoding. |
| loss_cls (obj:`mmcv.ConfigDict`|dict): Config of the |
| classification loss. Default `CrossEntropyLoss`. |
| loss_bbox (obj:`mmcv.ConfigDict`|dict): Config of the |
| regression loss. Default `L1Loss`. |
| loss_iou (obj:`mmcv.ConfigDict`|dict): Config of the |
| regression iou loss. Default `GIoULoss`. |
| tran_cfg (obj:`mmcv.ConfigDict`|dict): Training config of |
| transformer head. |
| test_cfg (obj:`mmcv.ConfigDict`|dict): Testing config of |
| transformer head. |
| init_cfg (dict or list[dict], optional): Initialization config dict. |
| Default: None |
| """ |
| _version = 2 |
|
|
| def __init__(self, |
| num_classes, |
| in_channels=256, |
| stride=[16], |
| embed_dims=256, |
| num_query=100, |
| num_reg_fcs=2, |
| memory_len=1024, |
| topk_proposals=256, |
| num_propagated=256, |
| with_dn=True, |
| with_ego_pos=True, |
| match_with_velo=True, |
| match_costs=None, |
| transformer=None, |
| sync_cls_avg_factor=False, |
| code_weights=None, |
| bbox_coder=None, |
| loss_cls=dict( |
| type='CrossEntropyLoss', |
| bg_cls_weight=0.1, |
| use_sigmoid=False, |
| loss_weight=1.0, |
| class_weight=1.0), |
| loss_bbox=dict(type='L1Loss', loss_weight=5.0), |
| loss_iou=dict(type='GIoULoss', loss_weight=2.0), |
| train_cfg=dict( |
| assigner=dict( |
| type='HungarianAssigner3D', |
| cls_cost=dict(type='ClassificationCost', weight=1.), |
| reg_cost=dict(type='BBoxL1Cost', weight=5.0), |
| iou_cost=dict( |
| type='IoUCost', iou_mode='giou', weight=2.0)),), |
| test_cfg=dict(max_per_img=100), |
| scalar = 5, |
| noise_scale = 0.4, |
| noise_trans = 0.0, |
| dn_weight = 1.0, |
| split = 0.5, |
| routing_mode='off', |
| routing_keep_ratio=1.0, |
| routing_keep_ratio_stopped=None, |
| routing_keep_ratio_moderate=None, |
| routing_keep_ratio_fast=None, |
| routing_stopped_speed=0.5, |
| routing_high_speed=12.0, |
| routing_log_stats=False, |
| routing_router=None, |
| init_cfg=None, |
| normedlinear=False, |
| **kwargs): |
| |
| |
| |
| if 'code_size' in kwargs: |
| self.code_size = kwargs['code_size'] |
| else: |
| self.code_size = 10 |
| if code_weights is not None: |
| self.code_weights = code_weights |
| else: |
| self.code_weights = [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 0.2, 0.2] |
|
|
| self.code_weights = self.code_weights[:self.code_size] |
|
|
| if match_costs is not None: |
| self.match_costs = match_costs |
| else: |
| self.match_costs = self.code_weights |
| |
| self.bg_cls_weight = 0 |
| self.sync_cls_avg_factor = sync_cls_avg_factor |
| class_weight = loss_cls.get('class_weight', None) |
| if class_weight is not None and (self.__class__ is SparseHead): |
| assert isinstance(class_weight, float), 'Expected ' \ |
| 'class_weight to have type float. Found ' \ |
| f'{type(class_weight)}.' |
| |
| |
| bg_cls_weight = loss_cls.get('bg_cls_weight', class_weight) |
| assert isinstance(bg_cls_weight, float), 'Expected ' \ |
| 'bg_cls_weight to have type float. Found ' \ |
| f'{type(bg_cls_weight)}.' |
| class_weight = torch.ones(num_classes + 1) * class_weight |
| |
| class_weight[num_classes] = bg_cls_weight |
| loss_cls.update({'class_weight': class_weight}) |
| if 'bg_cls_weight' in loss_cls: |
| loss_cls.pop('bg_cls_weight') |
| self.bg_cls_weight = bg_cls_weight |
|
|
| if train_cfg: |
| assert 'assigner' in train_cfg, 'assigner should be provided '\ |
| 'when train_cfg is set.' |
| assigner = train_cfg['assigner'] |
|
|
|
|
| self.assigner = build_assigner(assigner) |
| |
| sampler_cfg = dict(type='PseudoSampler') |
| self.sampler = build_sampler(sampler_cfg, context=self) |
|
|
| self.num_query = num_query |
| self.num_classes = num_classes |
| self.in_channels = in_channels |
| self.memory_len = memory_len |
| self.topk_proposals = topk_proposals |
| self.num_propagated = num_propagated |
| self.with_dn = with_dn |
| self.with_ego_pos = with_ego_pos |
| self.match_with_velo = match_with_velo |
| self.num_reg_fcs = num_reg_fcs |
| self.train_cfg = train_cfg |
| self.test_cfg = test_cfg |
| self.fp16_enabled = False |
| self.embed_dims = embed_dims |
| self.with_dn = with_dn |
| self.stride=stride |
|
|
| self.scalar = scalar |
| self.bbox_noise_scale = noise_scale |
| self.bbox_noise_trans = noise_trans |
| self.dn_weight = dn_weight |
| self.split = split |
| self.routing_mode = routing_mode |
| self.routing_keep_ratio = float(routing_keep_ratio) |
| self.routing_keep_ratio_stopped = routing_keep_ratio_stopped |
| self.routing_keep_ratio_moderate = routing_keep_ratio_moderate |
| self.routing_keep_ratio_fast = routing_keep_ratio_fast |
| self.routing_stopped_speed = float(routing_stopped_speed) |
| self.routing_high_speed = float(routing_high_speed) |
| self.routing_log_stats = routing_log_stats |
| self.routing_router_cfg = copy.deepcopy(routing_router) |
| self.latest_routing_metadata = None |
|
|
| self.act_cfg = transformer.get('act_cfg', |
| dict(type='ReLU', inplace=True)) |
| self.num_pred = 6 |
| self.normedlinear = normedlinear |
| super(SparseHead, self).__init__(num_classes, in_channels, init_cfg = init_cfg) |
|
|
| self.loss_cls = build_loss(loss_cls) |
| self.loss_bbox = build_loss(loss_bbox) |
| self.loss_iou = build_loss(loss_iou) |
|
|
| if self.loss_cls.use_sigmoid: |
| self.cls_out_channels = num_classes |
| else: |
| self.cls_out_channels = num_classes + 1 |
|
|
| self.transformer = build_transformer(transformer) |
|
|
| self.code_weights = nn.Parameter(torch.tensor( |
| self.code_weights), requires_grad=False) |
|
|
| self.match_costs = nn.Parameter(torch.tensor( |
| self.match_costs), requires_grad=False) |
|
|
| self.bbox_coder = build_bbox_coder(bbox_coder) |
| self.routing_router = self._build_routing_router(self.routing_router_cfg) |
|
|
| self.pc_range = nn.Parameter(torch.tensor( |
| self.bbox_coder.pc_range), requires_grad=False) |
|
|
|
|
| self._init_layers() |
| self.reset_memory() |
|
|
| def _build_routing_router(self, routing_router): |
| if routing_router is None: |
| if self.routing_mode == 'learned_gumbel_ambient': |
| return GumbelTokenRouter( |
| embed_dims=self.embed_dims, |
| keep_ratio=self.routing_keep_ratio, |
| ) |
| return None |
|
|
| routing_router = copy.deepcopy(routing_router) |
| router_type = routing_router.pop('type', 'GumbelTokenRouter') |
| if router_type != 'GumbelTokenRouter': |
| raise ValueError(f'Unsupported routing router type: {router_type}') |
| routing_router.setdefault('embed_dims', self.embed_dims) |
| routing_router.setdefault('keep_ratio', self.routing_keep_ratio) |
| return GumbelTokenRouter(**routing_router) |
|
|
| def _apply_token_routing(self, feat_flatten, data): |
| if self.routing_mode in (None, 'off'): |
| self.latest_routing_metadata = None |
| return feat_flatten |
|
|
| if self.routing_mode not in { |
| 'zero', 'hard_drop', 'ambient_pool', 'learned_gumbel_ambient' |
| }: |
| raise ValueError( |
| "routing_mode must be 'off', 'zero', 'hard_drop', or " |
| f"'ambient_pool', or 'learned_gumbel_ambient', got " |
| f"{self.routing_mode}" |
| ) |
|
|
| scores = feat_flatten.detach().float().norm(dim=-1) |
| keep_mask, target_ratio = make_budget_mask( |
| scores, |
| self.routing_keep_ratio, |
| ego_speed=data.get('ego_speed', None), |
| stopped_keep_ratio=self.routing_keep_ratio_stopped, |
| moderate_keep_ratio=self.routing_keep_ratio_moderate, |
| fast_keep_ratio=self.routing_keep_ratio_fast, |
| stopped_speed=self.routing_stopped_speed, |
| high_speed=self.routing_high_speed, |
| ) |
| if self.routing_mode == 'learned_gumbel_ambient': |
| routed, metadata = self.routing_router(feat_flatten, keep_ratio=target_ratio) |
| else: |
| routed, metadata = apply_dense_budget_mask( |
| feat_flatten, |
| keep_mask, |
| mode=self.routing_mode, |
| ) |
| metadata.update( |
| routing_target_keep_ratio=target_ratio.detach().cpu().tolist(), |
| routing_preserves_dense_geometry=True, |
| ) |
| self.latest_routing_metadata = metadata |
| if self.routing_log_stats: |
| print( |
| "routing_metadata: " |
| f"mode={metadata['routing_mode']} " |
| f"keep={metadata['routing_keep_count']}/" |
| f"{metadata['routing_num_tokens']}" |
| ) |
| return routed |
|
|
| def _init_layers(self): |
| """Initialize layers of the transformer head.""" |
| cls_branch = [] |
| for _ in range(self.num_reg_fcs): |
| cls_branch.append(Linear(self.embed_dims, self.embed_dims)) |
| cls_branch.append(nn.LayerNorm(self.embed_dims)) |
| cls_branch.append(nn.ReLU(inplace=True)) |
| if self.normedlinear: |
| cls_branch.append(NormedLinear(self.embed_dims, self.cls_out_channels)) |
| else: |
| cls_branch.append(Linear(self.embed_dims, self.cls_out_channels)) |
| fc_cls = nn.Sequential(*cls_branch) |
|
|
| reg_branch = [] |
| for _ in range(self.num_reg_fcs): |
| reg_branch.append(Linear(self.embed_dims, self.embed_dims)) |
| reg_branch.append(nn.ReLU()) |
| reg_branch.append(Linear(self.embed_dims, self.code_size)) |
| reg_branch = nn.Sequential(*reg_branch) |
|
|
| self.cls_branches = nn.ModuleList( |
| [fc_cls for _ in range(self.num_pred)]) |
| self.reg_branches = nn.ModuleList( |
| [reg_branch for _ in range(self.num_pred)]) |
|
|
| self.reference_points = nn.Embedding(self.num_query, 3) |
| if self.num_propagated > 0: |
| self.pseudo_reference_points = nn.Embedding(self.num_propagated, 3) |
|
|
|
|
| self.query_embedding = nn.Sequential( |
| nn.Linear(self.embed_dims*3//2, self.embed_dims), |
| nn.ReLU(), |
| nn.Linear(self.embed_dims, self.embed_dims), |
| ) |
|
|
| self.spatial_alignment = MLN(14, use_ln=False) |
|
|
| self.time_embedding = nn.Sequential( |
| nn.Linear(self.embed_dims, self.embed_dims), |
| nn.LayerNorm(self.embed_dims) |
| ) |
|
|
| |
| if self.with_ego_pos: |
| self.ego_pose_pe = MLN(180) |
| self.ego_pose_memory = MLN(180) |
|
|
| def temporal_alignment(self, query_pos, tgt, reference_points): |
| B = query_pos.size(0) |
|
|
| temp_reference_point = (self.memory_reference_point - self.pc_range[:3]) / (self.pc_range[3:6] - self.pc_range[0:3]) |
| temp_pos = self.query_embedding(pos2posemb3d(temp_reference_point)) |
| temp_memory = self.memory_embedding |
| rec_ego_pose = torch.eye(4, device=query_pos.device).unsqueeze(0).unsqueeze(0).repeat(B, query_pos.size(1), 1, 1) |
| |
| if self.with_ego_pos: |
| rec_ego_motion = torch.cat([torch.zeros_like(reference_points[...,:3]), rec_ego_pose[..., :3, :].flatten(-2)], dim=-1) |
| rec_ego_motion = nerf_positional_encoding(rec_ego_motion) |
| tgt = self.ego_pose_memory(tgt, rec_ego_motion) |
| query_pos = self.ego_pose_pe(query_pos, rec_ego_motion) |
| memory_ego_motion = torch.cat([self.memory_velo, self.memory_timestamp, self.memory_egopose[..., :3, :].flatten(-2)], dim=-1).float() |
| memory_ego_motion = nerf_positional_encoding(memory_ego_motion) |
| temp_pos = self.ego_pose_pe(temp_pos, memory_ego_motion) |
| temp_memory = self.ego_pose_memory(temp_memory, memory_ego_motion) |
|
|
| query_pos += self.time_embedding(pos2posemb1d(torch.zeros_like(reference_points[...,:1]))) |
| temp_pos += self.time_embedding(pos2posemb1d(self.memory_timestamp).float()) |
|
|
| if self.num_propagated > 0: |
| tgt = torch.cat([tgt, temp_memory[:, :self.num_propagated]], dim=1) |
| query_pos = torch.cat([query_pos, temp_pos[:, :self.num_propagated]], dim=1) |
| reference_points = torch.cat([reference_points, temp_reference_point[:, :self.num_propagated]], dim=1) |
| rec_ego_pose = torch.eye(4, device=query_pos.device).unsqueeze(0).unsqueeze(0).repeat(B, query_pos.shape[1]+self.num_propagated, 1, 1) |
| temp_memory = temp_memory[:, self.num_propagated:] |
| temp_pos = temp_pos[:, self.num_propagated:] |
| |
| return tgt, query_pos, reference_points, temp_memory, temp_pos, rec_ego_pose |
|
|
| def prepare_for_dn(self, batch_size, reference_points, img_metas): |
| if self.training and self.with_dn: |
| targets = [torch.cat((img_meta['gt_bboxes_3d']._data.gravity_center, img_meta['gt_bboxes_3d']._data.tensor[:, 3:]),dim=1) for img_meta in img_metas ] |
| labels = [img_meta['gt_labels_3d']._data for img_meta in img_metas ] |
| known = [(torch.ones_like(t)).cuda() for t in labels] |
| know_idx = known |
| unmask_bbox = unmask_label = torch.cat(known) |
| |
| known_num = [t.size(0) for t in targets] |
| |
| labels = torch.cat([t for t in labels]) |
| boxes = torch.cat([t for t in targets]) |
| batch_idx = torch.cat([torch.full((t.size(0), ), i) for i, t in enumerate(targets)]) |
| |
| known_indice = torch.nonzero(unmask_label + unmask_bbox) |
| known_indice = known_indice.view(-1) |
| |
| known_indice = known_indice.repeat(self.scalar, 1).view(-1) |
| known_labels = labels.repeat(self.scalar, 1).view(-1).long().to(reference_points.device) |
| known_bid = batch_idx.repeat(self.scalar, 1).view(-1) |
| known_bboxs = boxes.repeat(self.scalar, 1).to(reference_points.device) |
| known_bbox_center = known_bboxs[:, :3].clone() |
| known_bbox_scale = known_bboxs[:, 3:6].clone() |
|
|
| if self.bbox_noise_scale > 0: |
| diff = known_bbox_scale / 2 + self.bbox_noise_trans |
| rand_prob = torch.rand_like(known_bbox_center) * 2 - 1.0 |
| known_bbox_center += torch.mul(rand_prob, |
| diff) * self.bbox_noise_scale |
| known_bbox_center[..., 0:3] = (known_bbox_center[..., 0:3] - self.pc_range[0:3]) / (self.pc_range[3:6] - self.pc_range[0:3]) |
|
|
| known_bbox_center = known_bbox_center.clamp(min=0.0, max=1.0) |
| mask = torch.norm(rand_prob, 2, 1) > self.split |
| known_labels[mask] = self.num_classes |
| |
| single_pad = int(max(known_num)) |
| pad_size = int(single_pad * self.scalar) |
| padding_bbox = torch.zeros(pad_size, 3).to(reference_points.device) |
| padded_reference_points = torch.cat([padding_bbox, reference_points], dim=0).unsqueeze(0).repeat(batch_size, 1, 1) |
|
|
| if len(known_num): |
| map_known_indice = torch.cat([torch.tensor(range(num)) for num in known_num]) |
| map_known_indice = torch.cat([map_known_indice + single_pad * i for i in range(self.scalar)]).long() |
| if len(known_bid): |
| padded_reference_points[(known_bid.long(), map_known_indice)] = known_bbox_center.to(reference_points.device) |
|
|
| tgt_size = pad_size + self.num_query |
| attn_mask = torch.ones(tgt_size, tgt_size).to(reference_points.device) < 0 |
| |
| attn_mask[pad_size:, :pad_size] = True |
| |
| for i in range(self.scalar): |
| if i == 0: |
| attn_mask[single_pad * i:single_pad * (i + 1), single_pad * (i + 1):pad_size] = True |
| if i == self.scalar - 1: |
| attn_mask[single_pad * i:single_pad * (i + 1), :single_pad * i] = True |
| else: |
| attn_mask[single_pad * i:single_pad * (i + 1), single_pad * (i + 1):pad_size] = True |
| attn_mask[single_pad * i:single_pad * (i + 1), :single_pad * i] = True |
| |
| |
| query_size = pad_size + self.num_query + self.num_propagated |
| tgt_size = pad_size + self.num_query + self.memory_len |
| temporal_attn_mask = torch.ones(query_size, tgt_size).to(reference_points.device) < 0 |
| temporal_attn_mask[:attn_mask.size(0), :attn_mask.size(1)] = attn_mask |
| temporal_attn_mask[pad_size:, :pad_size] = True |
| attn_mask = temporal_attn_mask |
|
|
| mask_dict = { |
| 'known_indice': torch.as_tensor(known_indice).long(), |
| 'batch_idx': torch.as_tensor(batch_idx).long(), |
| 'map_known_indice': torch.as_tensor(map_known_indice).long(), |
| 'known_lbs_bboxes': (known_labels, known_bboxs), |
| 'know_idx': know_idx, |
| 'pad_size': pad_size |
| } |
| |
| else: |
| padded_reference_points = reference_points.unsqueeze(0).repeat(batch_size, 1, 1) |
| attn_mask = None |
| mask_dict = None |
|
|
| return padded_reference_points, attn_mask, mask_dict |
|
|
|
|
| def init_weights(self): |
| """Initialize weights of the transformer head.""" |
| |
| nn.init.uniform_(self.reference_points.weight.data, 0, 1) |
| if self.num_propagated > 0: |
| nn.init.uniform_(self.pseudo_reference_points.weight.data, 0, 1) |
| self.pseudo_reference_points.weight.requires_grad = False |
| self.transformer.init_weights() |
| if self.loss_cls.use_sigmoid: |
| bias_init = bias_init_with_prob(0.01) |
| for m in self.cls_branches: |
| nn.init.constant_(m[-1].bias, bias_init) |
|
|
|
|
| def reset_memory(self): |
| self.memory_embedding = None |
| self.memory_reference_point = None |
| self.memory_timestamp = None |
| self.memory_egopose = None |
| self.memory_velo = None |
|
|
| def pre_update_memory(self, data): |
| x = data['prev_exists'] |
| B = x.size(0) |
| |
| if self.memory_embedding is None: |
| self.memory_embedding = x.new_zeros(B, self.memory_len, self.embed_dims) |
| self.memory_reference_point = x.new_zeros(B, self.memory_len, 3) |
| self.memory_timestamp = x.new_zeros(B, self.memory_len, 1) |
| self.memory_egopose = x.new_zeros(B, self.memory_len, 4, 4) |
| self.memory_velo = x.new_zeros(B, self.memory_len, 2) |
| else: |
| self.memory_timestamp += data['timestamp'].unsqueeze(-1).unsqueeze(-1) |
| self.memory_egopose = data['ego_pose_inv'].unsqueeze(1) @ self.memory_egopose |
| self.memory_reference_point = transform_reference_points(self.memory_reference_point, data['ego_pose_inv'], reverse=False) |
| self.memory_timestamp = memory_refresh(self.memory_timestamp[:, :self.memory_len], x) |
| self.memory_reference_point = memory_refresh(self.memory_reference_point[:, :self.memory_len], x) |
| self.memory_embedding = memory_refresh(self.memory_embedding[:, :self.memory_len], x) |
| self.memory_egopose = memory_refresh(self.memory_egopose[:, :self.memory_len], x) |
| self.memory_velo = memory_refresh(self.memory_velo[:, :self.memory_len], x) |
| |
| |
| if self.num_propagated > 0: |
| pseudo_reference_points = self.pseudo_reference_points.weight * (self.pc_range[3:6] - self.pc_range[0:3]) + self.pc_range[0:3] |
| self.memory_reference_point[:, :self.num_propagated] = self.memory_reference_point[:, :self.num_propagated] + (1 - x).view(B, 1, 1) * pseudo_reference_points |
| self.memory_egopose[:, :self.num_propagated] = self.memory_egopose[:, :self.num_propagated] + (1 - x).view(B, 1, 1, 1) * torch.eye(4, device=x.device) |
|
|
| def post_update_memory(self, data, rec_ego_pose, all_cls_scores, all_bbox_preds, outs_dec, mask_dict): |
| if self.training and mask_dict and mask_dict['pad_size'] > 0: |
| rec_reference_points = all_bbox_preds[:, :, mask_dict['pad_size']:, :3][-1] |
| rec_velo = all_bbox_preds[:, :, mask_dict['pad_size']:, -2:][-1] |
| rec_memory = outs_dec[:, :, mask_dict['pad_size']:, :][-1] |
| rec_score = all_cls_scores[:, :, mask_dict['pad_size']:, :][-1].sigmoid().topk(1, dim=-1).values[..., 0:1] |
| rec_timestamp = torch.zeros_like(rec_score, dtype=torch.float64) |
| else: |
| rec_reference_points = all_bbox_preds[..., :3][-1] |
| rec_velo = all_bbox_preds[..., -2:][-1] |
| rec_memory = outs_dec[-1] |
| rec_score = all_cls_scores[-1].sigmoid().topk(1, dim=-1).values[..., 0:1] |
| rec_timestamp = torch.zeros_like(rec_score, dtype=torch.float64) |
| |
| |
| _, topk_indexes = torch.topk(rec_score, self.topk_proposals, dim=1) |
| rec_timestamp = topk_gather(rec_timestamp, topk_indexes) |
| rec_reference_points = topk_gather(rec_reference_points, topk_indexes).detach() |
| rec_memory = topk_gather(rec_memory, topk_indexes).detach() |
| rec_ego_pose = topk_gather(rec_ego_pose, topk_indexes) |
| rec_velo = topk_gather(rec_velo, topk_indexes).detach() |
|
|
| self.memory_embedding = torch.cat([rec_memory, self.memory_embedding], dim=1) |
| self.memory_timestamp = torch.cat([rec_timestamp, self.memory_timestamp], dim=1) |
| self.memory_egopose= torch.cat([rec_ego_pose, self.memory_egopose], dim=1) |
| self.memory_reference_point = torch.cat([rec_reference_points, self.memory_reference_point], dim=1) |
| self.memory_velo = torch.cat([rec_velo, self.memory_velo], dim=1) |
| self.memory_reference_point = transform_reference_points(self.memory_reference_point, data['ego_pose'], reverse=False) |
| self.memory_timestamp -= data['timestamp'].unsqueeze(-1).unsqueeze(-1) |
| self.memory_egopose = data['ego_pose'].unsqueeze(1) @ self.memory_egopose |
| |
| def forward(self, img_metas, **data): |
| """Forward function. |
| Args: |
| mlvl_feats (tuple[Tensor]): Features from the upstream |
| network, each is a 5D-tensor with shape |
| (B, N, C, H, W). |
| Returns: |
| all_cls_scores (Tensor): Outputs from the classification head, \ |
| shape [nb_dec, bs, num_query, cls_out_channels]. Note \ |
| cls_out_channels should includes background. |
| all_bbox_preds (Tensor): Sigmoid outputs from the regression \ |
| head with normalized coordinate format (cx, cy, w, l, cz, h, theta, vx, vy). \ |
| Shape [nb_dec, bs, num_query, 9]. |
| """ |
| self.pre_update_memory(data) |
| mlvl_feats = data['img_feats'] |
| B = mlvl_feats[0].size(0) |
|
|
| reference_points = self.reference_points.weight |
| dtype = reference_points.dtype |
| intrinsics = data['intrinsics'] / 1e3 |
| extrinsics = data['extrinsics'][..., :3, :] |
| mln_input = torch.cat([intrinsics[..., 0,0:1], intrinsics[..., 1,1:2], extrinsics.flatten(-2)], dim=-1) |
| mln_input = mln_input.flatten(0, 1).unsqueeze(1) |
| feat_flatten = [] |
| spatial_flatten = [] |
| for i in range(len(mlvl_feats)): |
| B, N, C, H, W = mlvl_feats[i].shape |
| mlvl_feat = mlvl_feats[i].reshape(B * N, C, -1).transpose(1, 2) |
| mlvl_feat = self.spatial_alignment(mlvl_feat, mln_input) |
| feat_flatten.append(mlvl_feat.to(dtype)) |
| spatial_flatten.append((H, W)) |
| feat_flatten = torch.cat(feat_flatten, dim=1) |
| feat_flatten = self._apply_token_routing(feat_flatten, data) |
| spatial_flatten = torch.as_tensor(spatial_flatten, dtype=torch.long, device=mlvl_feats[0].device) |
| level_start_index = torch.cat((spatial_flatten.new_zeros((1, )), spatial_flatten.prod(1).cumsum(0)[:-1])) |
| reference_points, attn_mask, mask_dict = self.prepare_for_dn(B, reference_points, img_metas) |
| query_pos = self.query_embedding(pos2posemb3d(reference_points)) |
| tgt = torch.zeros_like(query_pos) |
|
|
| |
| tgt, query_pos, reference_points, temp_memory, temp_pos, rec_ego_pose = self.temporal_alignment(query_pos, tgt, reference_points) |
|
|
| outs_dec = self.transformer(tgt, query_pos, feat_flatten, spatial_flatten, level_start_index, temp_memory, |
| temp_pos, attn_mask, reference_points, self.pc_range, data, img_metas) |
|
|
| outs_dec = torch.nan_to_num(outs_dec) |
| outputs_classes = [] |
| outputs_coords = [] |
| for lvl in range(outs_dec.shape[0]): |
| reference = inverse_sigmoid(reference_points.clone()) |
| assert reference.shape[-1] == 3 |
| outputs_class = self.cls_branches[lvl](outs_dec[lvl]) |
| tmp = self.reg_branches[lvl](outs_dec[lvl]) |
|
|
| tmp[..., 0:3] += reference[..., 0:3] |
| tmp[..., 0:3] = tmp[..., 0:3].sigmoid() |
|
|
| outputs_coord = tmp |
| outputs_classes.append(outputs_class) |
| outputs_coords.append(outputs_coord) |
|
|
| all_cls_scores = torch.stack(outputs_classes) |
| all_bbox_preds = torch.stack(outputs_coords) |
| all_bbox_preds[..., 0:3] = (all_bbox_preds[..., 0:3] * (self.pc_range[3:6] - self.pc_range[0:3]) + self.pc_range[0:3]) |
| |
| |
| self.post_update_memory(data, rec_ego_pose, all_cls_scores, all_bbox_preds, outs_dec, mask_dict) |
|
|
| if mask_dict and mask_dict['pad_size'] > 0: |
| output_known_class = all_cls_scores[:, :, :mask_dict['pad_size'], :] |
| output_known_coord = all_bbox_preds[:, :, :mask_dict['pad_size'], :] |
| outputs_class = all_cls_scores[:, :, mask_dict['pad_size']:, :] |
| outputs_coord = all_bbox_preds[:, :, mask_dict['pad_size']:, :] |
| mask_dict['output_known_lbs_bboxes']=(output_known_class, output_known_coord) |
| outs = { |
| 'all_cls_scores': outputs_class, |
| 'all_bbox_preds': outputs_coord, |
| 'dn_mask_dict':mask_dict, |
|
|
| } |
| else: |
| outs = { |
| 'all_cls_scores': all_cls_scores, |
| 'all_bbox_preds': all_bbox_preds, |
| 'dn_mask_dict':None, |
| } |
|
|
| return outs |
|
|
|
|
| def prepare_for_loss(self, mask_dict): |
| """ |
| prepare dn components to calculate loss |
| Args: |
| mask_dict: a dict that contains dn information |
| """ |
| output_known_class, output_known_coord = mask_dict['output_known_lbs_bboxes'] |
| known_labels, known_bboxs = mask_dict['known_lbs_bboxes'] |
| map_known_indice = mask_dict['map_known_indice'].long() |
| known_indice = mask_dict['known_indice'].long().cpu() |
| batch_idx = mask_dict['batch_idx'].long() |
| bid = batch_idx[known_indice] |
| if len(output_known_class) > 0: |
| output_known_class = output_known_class.permute(1, 2, 0, 3)[(bid, map_known_indice)].permute(1, 0, 2) |
| output_known_coord = output_known_coord.permute(1, 2, 0, 3)[(bid, map_known_indice)].permute(1, 0, 2) |
| num_tgt = known_indice.numel() |
| return known_labels, known_bboxs, output_known_class, output_known_coord, num_tgt |
|
|
|
|
| def _get_target_single(self, |
| cls_score, |
| bbox_pred, |
| gt_labels, |
| gt_bboxes, |
| gt_bboxes_ignore=None): |
| """"Compute regression and classification targets for one image. |
| Outputs from a single decoder layer of a single feature level are used. |
| Args: |
| cls_score (Tensor): Box score logits from a single decoder layer |
| for one image. Shape [num_query, cls_out_channels]. |
| bbox_pred (Tensor): Sigmoid outputs from a single decoder layer |
| for one image, with normalized coordinate (cx, cy, w, h) and |
| shape [num_query, 4]. |
| gt_bboxes (Tensor): Ground truth bboxes for one image with |
| shape (num_gts, 4) in [tl_x, tl_y, br_x, br_y] format. |
| gt_labels (Tensor): Ground truth class indexes for one image |
| with shape (num_gts, ). |
| gt_bboxes_ignore (Tensor, optional): Bounding boxes |
| which can be ignored. Default None. |
| Returns: |
| tuple[Tensor]: a tuple containing the following for one image. |
| - labels (Tensor): Labels of each image. |
| - label_weights (Tensor]): Label weights of each image. |
| - bbox_targets (Tensor): BBox targets of each image. |
| - bbox_weights (Tensor): BBox weights of each image. |
| - pos_inds (Tensor): Sampled positive indexes for each image. |
| - neg_inds (Tensor): Sampled negative indexes for each image. |
| """ |
|
|
| num_bboxes = bbox_pred.size(0) |
| |
|
|
| assign_result = self.assigner.assign(bbox_pred, cls_score, gt_bboxes, |
| gt_labels, gt_bboxes_ignore, self.match_costs, self.match_with_velo) |
| sampling_result = self.sampler.sample(assign_result, bbox_pred, |
| gt_bboxes) |
| pos_inds = sampling_result.pos_inds |
| neg_inds = sampling_result.neg_inds |
|
|
| |
| labels = gt_bboxes.new_full((num_bboxes, ), |
| self.num_classes, |
| dtype=torch.long) |
| label_weights = gt_bboxes.new_ones(num_bboxes) |
|
|
| |
| code_size = gt_bboxes.size(1) |
| bbox_targets = torch.zeros_like(bbox_pred)[..., :code_size] |
| bbox_weights = torch.zeros_like(bbox_pred) |
| |
| |
| if sampling_result.num_gts > 0: |
| bbox_targets[pos_inds] = sampling_result.pos_gt_bboxes |
| bbox_weights[pos_inds] = 1.0 |
| labels[pos_inds] = gt_labels[sampling_result.pos_assigned_gt_inds] |
| return (labels, label_weights, bbox_targets, bbox_weights, |
| pos_inds, neg_inds) |
|
|
| def get_targets(self, |
| cls_scores_list, |
| bbox_preds_list, |
| gt_bboxes_list, |
| gt_labels_list, |
| gt_bboxes_ignore_list=None): |
| """"Compute regression and classification targets for a batch image. |
| Outputs from a single decoder layer of a single feature level are used. |
| Args: |
| cls_scores_list (list[Tensor]): Box score logits from a single |
| decoder layer for each image with shape [num_query, |
| cls_out_channels]. |
| bbox_preds_list (list[Tensor]): Sigmoid outputs from a single |
| decoder layer for each image, with normalized coordinate |
| (cx, cy, w, h) and shape [num_query, 4]. |
| gt_bboxes_list (list[Tensor]): Ground truth bboxes for each image |
| with shape (num_gts, 4) in [tl_x, tl_y, br_x, br_y] format. |
| gt_labels_list (list[Tensor]): Ground truth class indexes for each |
| image with shape (num_gts, ). |
| gt_bboxes_ignore_list (list[Tensor], optional): Bounding |
| boxes which can be ignored for each image. Default None. |
| Returns: |
| tuple: a tuple containing the following targets. |
| - labels_list (list[Tensor]): Labels for all images. |
| - label_weights_list (list[Tensor]): Label weights for all \ |
| images. |
| - bbox_targets_list (list[Tensor]): BBox targets for all \ |
| images. |
| - bbox_weights_list (list[Tensor]): BBox weights for all \ |
| images. |
| - num_total_pos (int): Number of positive samples in all \ |
| images. |
| - num_total_neg (int): Number of negative samples in all \ |
| images. |
| """ |
| assert gt_bboxes_ignore_list is None, \ |
| 'Only supports for gt_bboxes_ignore setting to None.' |
| num_imgs = len(cls_scores_list) |
| gt_bboxes_ignore_list = [ |
| gt_bboxes_ignore_list for _ in range(num_imgs) |
| ] |
|
|
| (labels_list, label_weights_list, bbox_targets_list, |
| bbox_weights_list, pos_inds_list, neg_inds_list) = multi_apply( |
| self._get_target_single, cls_scores_list, bbox_preds_list, |
| gt_labels_list, gt_bboxes_list, gt_bboxes_ignore_list) |
| num_total_pos = sum((inds.numel() for inds in pos_inds_list)) |
| num_total_neg = sum((inds.numel() for inds in neg_inds_list)) |
| return (labels_list, label_weights_list, bbox_targets_list, |
| bbox_weights_list, num_total_pos, num_total_neg) |
|
|
| def loss_single(self, |
| cls_scores, |
| bbox_preds, |
| gt_bboxes_list, |
| gt_labels_list, |
| gt_bboxes_ignore_list=None): |
| """"Loss function for outputs from a single decoder layer of a single |
| feature level. |
| Args: |
| cls_scores (Tensor): Box score logits from a single decoder layer |
| for all images. Shape [bs, num_query, cls_out_channels]. |
| bbox_preds (Tensor): Sigmoid outputs from a single decoder layer |
| for all images, with normalized coordinate (cx, cy, w, h) and |
| shape [bs, num_query, 4]. |
| gt_bboxes_list (list[Tensor]): Ground truth bboxes for each image |
| with shape (num_gts, 4) in [tl_x, tl_y, br_x, br_y] format. |
| gt_labels_list (list[Tensor]): Ground truth class indexes for each |
| image with shape (num_gts, ). |
| gt_bboxes_ignore_list (list[Tensor], optional): Bounding |
| boxes which can be ignored for each image. Default None. |
| Returns: |
| dict[str, Tensor]: A dictionary of loss components for outputs from |
| a single decoder layer. |
| """ |
| num_imgs = cls_scores.size(0) |
| cls_scores_list = [cls_scores[i] for i in range(num_imgs)] |
| bbox_preds_list = [bbox_preds[i] for i in range(num_imgs)] |
| cls_reg_targets = self.get_targets(cls_scores_list, bbox_preds_list, |
| gt_bboxes_list, gt_labels_list, |
| gt_bboxes_ignore_list) |
| (labels_list, label_weights_list, bbox_targets_list, bbox_weights_list, |
| num_total_pos, num_total_neg) = cls_reg_targets |
| labels = torch.cat(labels_list, 0) |
| label_weights = torch.cat(label_weights_list, 0) |
| bbox_targets = torch.cat(bbox_targets_list, 0) |
| bbox_weights = torch.cat(bbox_weights_list, 0) |
|
|
| |
| cls_scores = cls_scores.reshape(-1, self.cls_out_channels) |
| |
| cls_avg_factor = num_total_pos * 1.0 + \ |
| num_total_neg * self.bg_cls_weight |
| if self.sync_cls_avg_factor: |
| cls_avg_factor = reduce_mean( |
| cls_scores.new_tensor([cls_avg_factor])) |
|
|
| cls_avg_factor = max(cls_avg_factor, 1) |
| loss_cls = self.loss_cls( |
| cls_scores, labels, label_weights, avg_factor=cls_avg_factor) |
|
|
| |
| |
| num_total_pos = loss_cls.new_tensor([num_total_pos]) |
| num_total_pos = torch.clamp(reduce_mean(num_total_pos), min=1).item() |
|
|
| |
| bbox_preds = bbox_preds.reshape(-1, bbox_preds.size(-1)) |
| normalized_bbox_targets = normalize_bbox(bbox_targets, self.pc_range) |
| isnotnan = torch.isfinite(normalized_bbox_targets).all(dim=-1) |
| bbox_weights = bbox_weights * self.code_weights |
|
|
| loss_bbox = self.loss_bbox( |
| bbox_preds[isnotnan, :10], normalized_bbox_targets[isnotnan, :10], bbox_weights[isnotnan, :10], avg_factor=num_total_pos) |
|
|
| loss_cls = torch.nan_to_num(loss_cls) |
| loss_bbox = torch.nan_to_num(loss_bbox) |
| return loss_cls, loss_bbox |
|
|
| |
| def dn_loss_single(self, |
| cls_scores, |
| bbox_preds, |
| known_bboxs, |
| known_labels, |
| num_total_pos=None): |
| """"Loss function for outputs from a single decoder layer of a single |
| feature level. |
| Args: |
| cls_scores (Tensor): Box score logits from a single decoder layer |
| for all images. Shape [bs, num_query, cls_out_channels]. |
| bbox_preds (Tensor): Sigmoid outputs from a single decoder layer |
| for all images, with normalized coordinate (cx, cy, w, h) and |
| shape [bs, num_query, 4]. |
| gt_bboxes_list (list[Tensor]): Ground truth bboxes for each image |
| with shape (num_gts, 4) in [tl_x, tl_y, br_x, br_y] format. |
| gt_labels_list (list[Tensor]): Ground truth class indexes for each |
| image with shape (num_gts, ). |
| gt_bboxes_ignore_list (list[Tensor], optional): Bounding |
| boxes which can be ignored for each image. Default None. |
| Returns: |
| dict[str, Tensor]: A dictionary of loss components for outputs from |
| a single decoder layer. |
| """ |
| |
| cls_scores = cls_scores.reshape(-1, self.cls_out_channels) |
| |
| cls_avg_factor = num_total_pos * 3.14159 / 6 * self.split * self.split * self.split |
| if self.sync_cls_avg_factor: |
| cls_avg_factor = reduce_mean( |
| cls_scores.new_tensor([cls_avg_factor])) |
| bbox_weights = torch.ones_like(bbox_preds) |
| label_weights = torch.ones_like(known_labels) |
| cls_avg_factor = max(cls_avg_factor, 1) |
| loss_cls = self.loss_cls( |
| cls_scores, known_labels.long(), label_weights, avg_factor=cls_avg_factor) |
|
|
| |
| |
| num_total_pos = loss_cls.new_tensor([num_total_pos]) |
| num_total_pos = torch.clamp(reduce_mean(num_total_pos), min=1).item() |
|
|
| |
| bbox_preds = bbox_preds.reshape(-1, bbox_preds.size(-1)) |
| normalized_bbox_targets = normalize_bbox(known_bboxs, self.pc_range) |
| isnotnan = torch.isfinite(normalized_bbox_targets).all(dim=-1) |
|
|
| bbox_weights = bbox_weights * self.code_weights |
|
|
| |
| loss_bbox = self.loss_bbox( |
| bbox_preds[isnotnan, :10], normalized_bbox_targets[isnotnan, :10], bbox_weights[isnotnan, :10], avg_factor=num_total_pos) |
|
|
| loss_cls = torch.nan_to_num(loss_cls) |
| loss_bbox = torch.nan_to_num(loss_bbox) |
| |
| return self.dn_weight * loss_cls, self.dn_weight * loss_bbox |
| |
| @force_fp32(apply_to=('preds_dicts')) |
| def loss(self, |
| gt_bboxes_list, |
| gt_labels_list, |
| preds_dicts, |
| gt_bboxes_ignore=None): |
| """"Loss function. |
| Args: |
| gt_bboxes_list (list[Tensor]): Ground truth bboxes for each image |
| with shape (num_gts, 4) in [tl_x, tl_y, br_x, br_y] format. |
| gt_labels_list (list[Tensor]): Ground truth class indexes for each |
| image with shape (num_gts, ). |
| preds_dicts: |
| all_cls_scores (Tensor): Classification score of all |
| decoder layers, has shape |
| [nb_dec, bs, num_query, cls_out_channels]. |
| all_bbox_preds (Tensor): Sigmoid regression |
| outputs of all decode layers. Each is a 4D-tensor with |
| normalized coordinate format (cx, cy, w, h) and shape |
| [nb_dec, bs, num_query, 4]. |
| enc_cls_scores (Tensor): Classification scores of |
| points on encode feature map , has shape |
| (N, h*w, num_classes). Only be passed when as_two_stage is |
| True, otherwise is None. |
| enc_bbox_preds (Tensor): Regression results of each points |
| on the encode feature map, has shape (N, h*w, 4). Only be |
| passed when as_two_stage is True, otherwise is None. |
| gt_bboxes_ignore (list[Tensor], optional): Bounding boxes |
| which can be ignored for each image. Default None. |
| Returns: |
| dict[str, Tensor]: A dictionary of loss components. |
| """ |
| assert gt_bboxes_ignore is None, \ |
| f'{self.__class__.__name__} only supports ' \ |
| f'for gt_bboxes_ignore setting to None.' |
|
|
| all_cls_scores = preds_dicts['all_cls_scores'] |
| all_bbox_preds = preds_dicts['all_bbox_preds'] |
|
|
| num_dec_layers = len(all_cls_scores) |
| device = gt_labels_list[0].device |
| gt_bboxes_list = [torch.cat( |
| (gt_bboxes.gravity_center, gt_bboxes.tensor[:, 3:]), |
| dim=1).to(device) for gt_bboxes in gt_bboxes_list] |
|
|
| all_gt_bboxes_list = [gt_bboxes_list for _ in range(num_dec_layers)] |
| all_gt_labels_list = [gt_labels_list for _ in range(num_dec_layers)] |
| all_gt_bboxes_ignore_list = [ |
| gt_bboxes_ignore for _ in range(num_dec_layers) |
| ] |
|
|
| losses_cls, losses_bbox = multi_apply( |
| self.loss_single, all_cls_scores, all_bbox_preds, |
| all_gt_bboxes_list, all_gt_labels_list, |
| all_gt_bboxes_ignore_list) |
|
|
| loss_dict = dict() |
|
|
| |
| |
| loss_dict['loss_cls'] = losses_cls[-1] |
| loss_dict['loss_bbox'] = losses_bbox[-1] |
|
|
| |
| num_dec_layer = 0 |
| for loss_cls_i, loss_bbox_i in zip(losses_cls[:-1], |
| losses_bbox[:-1]): |
| loss_dict[f'd{num_dec_layer}.loss_cls'] = loss_cls_i |
| loss_dict[f'd{num_dec_layer}.loss_bbox'] = loss_bbox_i |
| num_dec_layer += 1 |
| |
| if preds_dicts['dn_mask_dict'] is not None: |
| known_labels, known_bboxs, output_known_class, output_known_coord, num_tgt = self.prepare_for_loss(preds_dicts['dn_mask_dict']) |
| all_known_bboxs_list = [known_bboxs for _ in range(num_dec_layers)] |
| all_known_labels_list = [known_labels for _ in range(num_dec_layers)] |
| all_num_tgts_list = [ |
| num_tgt for _ in range(num_dec_layers) |
| ] |
| |
| dn_losses_cls, dn_losses_bbox = multi_apply( |
| self.dn_loss_single, output_known_class, output_known_coord, |
| all_known_bboxs_list, all_known_labels_list, |
| all_num_tgts_list) |
| loss_dict['dn_loss_cls'] = dn_losses_cls[-1] |
| loss_dict['dn_loss_bbox'] = dn_losses_bbox[-1] |
| num_dec_layer = 0 |
| for loss_cls_i, loss_bbox_i in zip(dn_losses_cls[:-1], |
| dn_losses_bbox[:-1]): |
| loss_dict[f'd{num_dec_layer}.dn_loss_cls'] = loss_cls_i |
| loss_dict[f'd{num_dec_layer}.dn_loss_bbox'] = loss_bbox_i |
| num_dec_layer += 1 |
| |
| elif self.with_dn: |
| dn_losses_cls, dn_losses_bbox = multi_apply( |
| self.loss_single, all_cls_scores, all_bbox_preds, |
| all_gt_bboxes_list, all_gt_labels_list, |
| all_gt_bboxes_ignore_list) |
| loss_dict['dn_loss_cls'] = dn_losses_cls[-1].detach() |
| loss_dict['dn_loss_bbox'] = dn_losses_bbox[-1].detach() |
| num_dec_layer = 0 |
| for loss_cls_i, loss_bbox_i in zip(dn_losses_cls[:-1], |
| dn_losses_bbox[:-1]): |
| loss_dict[f'd{num_dec_layer}.dn_loss_cls'] = loss_cls_i.detach() |
| loss_dict[f'd{num_dec_layer}.dn_loss_bbox'] = loss_bbox_i.detach() |
| num_dec_layer += 1 |
|
|
| return loss_dict |
|
|
|
|
| @force_fp32(apply_to=('preds_dicts')) |
| def get_bboxes(self, preds_dicts, img_metas, rescale=False): |
| """Generate bboxes from bbox head predictions. |
| Args: |
| preds_dicts (tuple[list[dict]]): Prediction results. |
| img_metas (list[dict]): Point cloud and image's meta info. |
| Returns: |
| list[dict]: Decoded bbox, scores and labels after nms. |
| """ |
| preds_dicts = self.bbox_coder.decode(preds_dicts) |
| num_samples = len(preds_dicts) |
|
|
| ret_list = [] |
| for i in range(num_samples): |
| preds = preds_dicts[i] |
| bboxes = preds['bboxes'] |
| bboxes[:, 2] = bboxes[:, 2] - bboxes[:, 5] * 0.5 |
| bboxes = img_metas[i]['box_type_3d'](bboxes, bboxes.size(-1)) |
| scores = preds['scores'] |
| labels = preds['labels'] |
| ret_list.append([bboxes, scores, labels]) |
| return ret_list |
|
|
| class MLN(nn.Module): |
| ''' |
| Args: |
| c_dim (int): dimension of latent code c |
| f_dim (int): feature dimension |
| ''' |
|
|
| def __init__(self, c_dim, f_dim=256, use_ln=True): |
| super().__init__() |
| self.c_dim = c_dim |
| self.f_dim = f_dim |
| self.use_ln = use_ln |
|
|
| self.reduce = nn.Sequential( |
| nn.Linear(c_dim, f_dim), |
| nn.ReLU(), |
| ) |
| self.gamma = nn.Linear(f_dim, f_dim) |
| self.beta = nn.Linear(f_dim, f_dim) |
| if self.use_ln: |
| self.ln = nn.LayerNorm(f_dim, elementwise_affine=False) |
| self.init_weight() |
|
|
| def init_weight(self): |
| nn.init.zeros_(self.gamma.weight) |
| nn.init.zeros_(self.beta.weight) |
| nn.init.ones_(self.gamma.bias) |
| nn.init.zeros_(self.beta.bias) |
|
|
| def forward(self, x, c): |
| if self.use_ln: |
| x = self.ln(x) |
| c = self.reduce(c) |
| gamma = self.gamma(c) |
| beta = self.beta(c) |
| out = gamma * x + beta |
|
|
| return out |
|
|