Source code for pimm.models.sonata.sonata_v1m1_base

"""
Sonata v1m1 Base

Author: Xiaoyang Wu (xiaoyang.wu.cs@gmail.com)
Please cite our work if the code is helpful to you.
"""

from itertools import chain
from packaging import version
from functools import partial
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
import torch_scatter
from timm.layers import trunc_normal_

import pointops
from pimm.models.utils.structure import Point
from pimm.models.builder import MODELS, build_model
from pimm.models.modules import PointModel, PointModule  # noqa: F401
from pimm.models.utils import offset2batch, offset2bincount, batch2offset
from pimm.utils.comm import get_world_size, all_gather  # noqa: F401
from pimm.utils.scheduler import CosineScheduler


class OnlineCluster(nn.Module):
    def __init__(
        self,
        in_channels,
        hidden_channels=4096,
        embed_channels=512,
        num_prototypes=4096,
    ):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(in_channels, hidden_channels),
            nn.GELU(),
            nn.Linear(hidden_channels, embed_channels),
        )
        self.apply(self._init_weights)
        if version.parse(torch.__version__) >= version.parse("2.1.0"):
            self.prototype = torch.nn.utils.parametrizations.weight_norm(
                nn.Linear(embed_channels, num_prototypes, bias=False)
            )
            self.prototype.parametrizations.weight.original0.data.fill_(1)
            self.prototype.parametrizations.weight.original0.requires_grad = False

        else:
            self.prototype = torch.nn.utils.weight_norm(
                nn.Linear(embed_channels, num_prototypes, bias=False)
            )
            self.prototype.weight_g.data.fill_(1)
            self.prototype.weight_g.requires_grad = False

    @staticmethod
    def _init_weights(m):
        if isinstance(m, nn.Linear):
            trunc_normal_(m.weight, std=0.02)
            if isinstance(m, nn.Linear) and m.bias is not None:
                nn.init.constant_(m.bias, 0)

    def forward(self, feat):
        feat = self.mlp(feat)
        eps = 1e-6 if feat.dtype == torch.float16 else 1e-12
        feat = nn.functional.normalize(feat, dim=-1, p=2, eps=eps)
        similarity = self.prototype(feat)
        return similarity


class RepresentationFusion(nn.Module):
    def __init__(
        self,
        in_channels,
        out_channels,
        hidden_channels=None,
        dropout=0.0,
        residual=True,
        output_norm=True,
    ):
        super().__init__()
        hidden_channels = hidden_channels or max(in_channels, out_channels * 2)
        self.residual = residual
        self.input_norm = nn.LayerNorm(in_channels)
        self.mlp = nn.Sequential(
            nn.Linear(in_channels, hidden_channels),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_channels, out_channels),
        )
        self.skip = (
            nn.Linear(in_channels, out_channels, bias=False)
            if residual
            else None
        )
        self.output_norm = nn.LayerNorm(out_channels) if output_norm else nn.Identity()
        self.apply(self._init_weights)

    @staticmethod
    def _init_weights(m):
        if isinstance(m, nn.Linear):
            trunc_normal_(m.weight, std=0.02)
            if m.bias is not None:
                nn.init.constant_(m.bias, 0)

    def forward(self, feat):
        feat_norm = self.input_norm(feat)
        out = self.mlp(feat_norm)
        if self.skip is not None:
            out = out + self.skip(feat_norm)
        return self.output_norm(out)


[docs] @MODELS.register_module("Sonata-v1m1") class Sonata(PointModel): def __init__( self, backbone, head_in_channels, head_hidden_channels=4096, head_embed_channels=512, head_num_prototypes=4096, teacher_custom=None, num_global_view=2, num_local_view=4, mask_size_start=0.1, mask_size_base=0.4, mask_size_warmup_ratio=0.05, mask_ratio_start=0.3, mask_ratio_base=0.7, mask_ratio_warmup_ratio=0.05, mask_jitter=None, mask_jitter_start=0.0, mask_jitter_base=0.01, mask_jitter_warmup_ratio=0.05, teacher_temp_start=0.04, teacher_temp_base=0.07, teacher_temp_warmup_ratio=0.05, student_temp=0.1, mask_loss_weight=2 / 8, roll_mask_loss_weight=2 / 8, unmask_loss_weight=4 / 8, momentum_base=0.996, momentum_final=1, match_max_k=8, match_max_r=0.08, up_cast_level=2, representation_fusion_channels=None, representation_fusion_hidden_channels=None, representation_fusion_dropout=0.0, representation_fusion_residual=True, representation_fusion_output_norm=True, ): super(Sonata, self).__init__() self.mask_loss_weight = mask_loss_weight self.roll_mask_loss_weight = roll_mask_loss_weight self.unmask_loss_weight = unmask_loss_weight self.num_global_view = num_global_view self.num_local_view = num_local_view # masking and scheduler self.mask_size = mask_size_start self.mask_size_start = mask_size_start self.mask_size_base = mask_size_base self.mask_size_warmup_ratio = mask_size_warmup_ratio self.mask_size_scheduler = None self.mask_ratio = mask_ratio_start self.mask_ratio_start = mask_ratio_start self.mask_ratio_base = mask_ratio_base self.mask_ratio_warmup_ratio = mask_ratio_warmup_ratio self.mask_ratio_scheduler = None self.mask_jitter = mask_jitter if mask_jitter is None: self.mask_jitter_start = mask_jitter_start self.mask_jitter_base = mask_jitter_base self.mask_jitter_warmup_ratio = mask_jitter_warmup_ratio self.mask_jitter_scheduler = None # temperature and scheduler self.teacher_temp = teacher_temp_start self.teacher_temp_start = teacher_temp_start self.teacher_temp_base = teacher_temp_base self.teacher_temp_warmup_ratio = teacher_temp_warmup_ratio self.teacher_temp_scheduler = None self.student_temp = student_temp # momentum and scheduler self.momentum = momentum_base self.momentum_base = momentum_base self.momentum_final = momentum_final self.momentum_scheduler = None # dynamic matching self.match_max_k = match_max_k self.match_max_r = match_max_r # up cast level self.up_cast_level = up_cast_level self.representation_fusion_enabled = representation_fusion_channels is not None self.representation_fusion_channels = representation_fusion_channels head_feature_channels = ( representation_fusion_channels if self.representation_fusion_enabled else head_in_channels ) # one of unmask, mask, roll mask loss enable assert unmask_loss_weight + mask_loss_weight + roll_mask_loss_weight > 0 # roll mask loss need more than one global view assert num_global_view > 1 or roll_mask_loss_weight == 0 # current roll mask only support two global views assert num_global_view == 1 or num_global_view == 2 student_model_dict = dict() teacher_model_dict = dict() if teacher_custom is None: teacher_custom = {} student_backbone = build_model(backbone) # turn off parameters like drop path for teacher model backbone.update(teacher_custom) teacher_backbone = build_model(backbone) student_model_dict["backbone"] = student_backbone teacher_model_dict["backbone"] = teacher_backbone if self.representation_fusion_enabled: student_model_dict["representation_fusion"] = RepresentationFusion( in_channels=head_in_channels, out_channels=representation_fusion_channels, hidden_channels=representation_fusion_hidden_channels, dropout=representation_fusion_dropout, residual=representation_fusion_residual, output_norm=representation_fusion_output_norm, ) teacher_model_dict["representation_fusion"] = RepresentationFusion( in_channels=head_in_channels, out_channels=representation_fusion_channels, hidden_channels=representation_fusion_hidden_channels, dropout=representation_fusion_dropout, residual=representation_fusion_residual, output_norm=representation_fusion_output_norm, ) head = partial( OnlineCluster, in_channels=head_feature_channels, hidden_channels=head_hidden_channels, embed_channels=head_embed_channels, num_prototypes=head_num_prototypes, ) if self.mask_loss_weight > 0 or self.roll_mask_loss_weight > 0: student_model_dict["mask_head"] = head() teacher_model_dict["mask_head"] = head() if self.unmask_loss_weight > 0: student_model_dict["unmask_head"] = head() teacher_model_dict["unmask_head"] = head() self.student = nn.ModuleDict(student_model_dict) self.teacher = nn.ModuleDict(teacher_model_dict) for k, v in self.student.items(): self.teacher[k].load_state_dict(self.student[k].state_dict()) for p in self.teacher.parameters(): p.requires_grad = False
[docs] def before_train(self): # make ModelHook after CheckPointLoader total_steps = self.trainer.cfg.scheduler.total_steps curr_step = getattr(self.trainer, "global_step", 0) or ( self.trainer.start_epoch * len(self.trainer.train_loader) ) # mask size scheduler self.mask_size_scheduler = CosineScheduler( start_value=self.mask_size_start, base_value=self.mask_size_base, final_value=self.mask_size_base, warmup_iters=int(total_steps * self.mask_size_warmup_ratio), total_iters=total_steps, ) self.mask_size_scheduler.iter = curr_step # mask ratio scheduler self.mask_ratio_scheduler = CosineScheduler( start_value=self.mask_ratio_start, base_value=self.mask_ratio_base, final_value=self.mask_ratio_base, warmup_iters=int(total_steps * self.mask_ratio_warmup_ratio), total_iters=total_steps, ) self.mask_ratio_scheduler.iter = curr_step # teacher temperature scheduler self.teacher_temp_scheduler = CosineScheduler( start_value=self.teacher_temp_start, base_value=self.teacher_temp_base, final_value=self.teacher_temp_base, warmup_iters=int(total_steps * self.teacher_temp_warmup_ratio), total_iters=total_steps, ) self.teacher_temp_scheduler.iter = curr_step # momentum scheduler self.momentum_scheduler = CosineScheduler( base_value=self.momentum_base, final_value=self.momentum_final, total_iters=total_steps, ) self.momentum_scheduler.iter = curr_step if self.mask_jitter is None: # mask jitter scheduler self.mask_jitter_scheduler = CosineScheduler( start_value=self.mask_jitter_start, base_value=self.mask_jitter_start, final_value=self.mask_jitter_base, warmup_iters=int(total_steps * self.mask_jitter_warmup_ratio), total_iters=total_steps, ) self.mask_jitter_scheduler.iter = curr_step
[docs] def before_step(self): # update parameters from schedulers self.mask_size = self.mask_size_scheduler.step() self.mask_ratio = self.mask_ratio_scheduler.step() self.teacher_temp = self.teacher_temp_scheduler.step() self.momentum = self.momentum_scheduler.step() if hasattr(self, "mask_jitter_scheduler"): self.mask_jitter = self.mask_jitter_scheduler.step() if self.trainer.writer is not None: self.trainer.writer.add_scalar( "params/mask_size", self.mask_size, self.mask_size_scheduler.iter, ) self.trainer.writer.add_scalar( "params/mask_ratio", self.mask_ratio, self.mask_ratio_scheduler.iter, ) self.trainer.writer.add_scalar( "params/teacher_temp", self.teacher_temp, self.teacher_temp_scheduler.iter, ) self.trainer.writer.add_scalar( "params/momentum", self.momentum, self.momentum_scheduler.iter, ) if hasattr(self, "mask_jitter_scheduler"): self.trainer.writer.add_scalar( "params/mask_jitter", self.mask_jitter, self.mask_jitter_scheduler.iter, )
[docs] def after_step(self): # EMA update teacher with torch.no_grad(): m = self.momentum student_param_list = list(self.student.parameters()) teacher_param_list = list(self.teacher.parameters()) torch._foreach_mul_(teacher_param_list, m) torch._foreach_add_(teacher_param_list, student_param_list, alpha=1 - m)
[docs] @staticmethod def sinkhorn_knopp(feat, temp, num_iter=3): feat = feat.float() q = torch.exp(feat / temp).t() k = q.shape[0] # number of prototypes # batch n and sum_q into single allreduce n_local = q.shape[1] sum_q_local = q.sum() if get_world_size() > 1: scalars = torch.stack([q.new_tensor(n_local), sum_q_local]) dist.all_reduce(scalars) n, sum_q = scalars[0], scalars[1] else: n, sum_q = q.new_tensor(n_local), sum_q_local q = q / sum_q for i in range(num_iter): # normalize each row: total weight per prototype must be 1/k q_row_sum = q.sum(dim=1, keepdim=True) if get_world_size() > 1: dist.all_reduce(q_row_sum) q = q / q_row_sum / k # normalize each column: total weight per sample must be 1/n q = q / q.sum(dim=0, keepdim=True) / n q *= n # the columns must sum to 1 so that Q is an assignment return q.t()
[docs] def generate_mask(self, coord, offset): batch = offset2batch(offset) mask_size = self.mask_size mask_ratio = self.mask_ratio # Grouping points with grid patch min_coord = torch_scatter.segment_coo(coord, batch, reduce="min") grid_coord = ((coord - min_coord[batch]) // mask_size).int() grid_coord = torch.cat([batch.unsqueeze(-1), grid_coord], dim=-1) unique, point_cluster, counts = torch.unique( grid_coord, dim=0, sorted=True, return_inverse=True, return_counts=True ) patch_num = unique.shape[0] mask_patch_num = int(patch_num * mask_ratio) patch_index = torch.randperm(patch_num, device=coord.device) mask_patch_index = patch_index[:mask_patch_num] point_mask = torch.isin(point_cluster, mask_patch_index) return point_mask, point_cluster
[docs] @torch.no_grad() def match_neighbour( self, view1_coord, view1_offset, view2_coord, view2_offset, ): index2, distance = pointops.knn_query( 1, view2_coord.float(), view2_offset.int(), view1_coord.float(), view1_offset.int(), ) index1 = torch.arange( index2.shape[0], device=index2.device, dtype=torch.long ).unsqueeze(-1) index = torch.cat([index1, index2], dim=-1)[distance.squeeze(-1) < self.match_max_r] return index
[docs] @torch.no_grad() def roll_point(self, point): n = self.num_global_view # [pc1, pc1', pc2, pc2'] -> [pc1', pc1, pc2', pc2], only support num_global_view == 2 bs = len(point.offset) // self.num_global_view data_dict = {} for key in point.keys(): if key in ["feat", "coord", "origin_coord", "batch"]: value = point[key].split(offset2bincount(point.offset).tolist()) value = chain(*[value[n * b : n * (b + 1)][::-1] for b in range(bs)]) if key == "batch": value = [torch.ones_like(v) * i for i, v in enumerate(value)] data_dict[key] = torch.cat(list(value), dim=0) return Point(data_dict)
[docs] def up_cast(self, point): for _ in range(self.up_cast_level): assert "pooling_parent" in point.keys() assert "pooling_inverse" in point.keys() parent = point.pop("pooling_parent") inverse = point.pop("pooling_inverse") parent.feat = torch.cat([parent.feat, point.feat[inverse]], dim=-1) point = parent return point
[docs] def upsample_to_original(self, point): while "pooling_parent" in point.keys(): parent = point.pop("pooling_parent") inverse = point.pop("pooling_inverse") parent.feat = point.feat[inverse] point = parent return point
[docs] def fuse_representation(self, point, model_dict): if self.representation_fusion_enabled: point.feat = model_dict["representation_fusion"](point.feat) return point
[docs] def forward(self, data_dict, return_point=False): if return_point: point = self.teacher.backbone(data_dict) point = self.up_cast(point) point = self.fuse_representation(point, self.teacher) if self.representation_fusion_enabled: point = self.upsample_to_original(point) return dict(point=point) # prepare global_point, mask_global_point, local_point with torch.no_grad(): # global_point & masking global_point = Point( feat=data_dict["global_feat"], coord=data_dict["global_coord"], origin_coord=data_dict["global_origin_coord"], offset=data_dict["global_offset"], grid_size=data_dict["grid_size"][0], ) global_mask, global_cluster = self.generate_mask( global_point.coord, global_point.offset ) mask_global_coord = global_point.coord.clone().detach() if self.mask_jitter is not None: mask_global_coord[global_mask] += torch.clip( torch.randn_like(mask_global_coord[global_mask]).mul( self.mask_jitter ), max=self.mask_jitter * 2, ) mask_global_point = Point( feat=data_dict["global_feat"], coord=mask_global_coord, origin_coord=data_dict["global_origin_coord"], mask=global_mask, offset=data_dict["global_offset"], grid_size=data_dict["grid_size"][0], ) # local point & matching local_point = Point( feat=data_dict["local_feat"], coord=data_dict["local_coord"], origin_coord=data_dict["local_origin_coord"], offset=data_dict["local_offset"], grid_size=data_dict["grid_size"][0], ) # create result dictionary for return result_dict = dict(loss=[]) # teacher backbone forward (shared with mask and unmask) global_point_ = self.teacher.backbone(global_point) global_point_ = self.up_cast(global_point_) global_point_ = self.fuse_representation(global_point_, self.teacher) global_feat = global_point_.feat if self.mask_loss_weight > 0 or self.roll_mask_loss_weight > 0: # teacher head forward with torch.no_grad(): global_point_.feat = self.teacher.mask_head(global_feat) # student forward mask_global_point_ = self.student.backbone(mask_global_point) mask_global_point_ = self.up_cast(mask_global_point_) mask_global_point_ = self.fuse_representation(mask_global_point_, self.student) mask_pred_sim = self.student.mask_head(mask_global_point_.feat) if self.mask_loss_weight > 0: with torch.no_grad(): match_index = self.match_neighbour( mask_global_point_.origin_coord, mask_global_point_.offset, global_point_.origin_coord, global_point_.offset, ) # teacher forward mask_target_sim = self.sinkhorn_knopp( global_point_.feat[match_index[:, 1]], self.teacher_temp, ) # loss mask_loss = -torch.sum( mask_target_sim * F.log_softmax( mask_pred_sim[match_index[:, 0]] / self.student_temp, dim=-1 ), dim=-1, ) mask_loss = torch_scatter.segment_coo( mask_loss, index=mask_global_point_.batch[match_index[:, 0]], reduce="mean", ).mean() result_dict["mask_loss"] = mask_loss result_dict["loss"].append(mask_loss * self.mask_loss_weight) if self.roll_mask_loss_weight > 0: roll_global_point_ = self.roll_point(global_point_) with torch.no_grad(): # match index for pred and roll target match_index = self.match_neighbour( mask_global_point_.origin_coord, mask_global_point_.offset, roll_global_point_.origin_coord, roll_global_point_.offset, ) # teacher forward roll_mask_target_sim = self.sinkhorn_knopp( roll_global_point_.feat[match_index[:, 1]], self.teacher_temp, ) roll_mask_loss = -torch.sum( roll_mask_target_sim * F.log_softmax( mask_pred_sim[match_index[:, 0]] / self.student_temp, dim=-1 ), dim=-1, ) roll_mask_loss = torch_scatter.segment_coo( roll_mask_loss, index=mask_global_point_.batch[match_index[:, 0]], reduce="mean", ).mean() result_dict["roll_mask_loss"] = roll_mask_loss result_dict["loss"].append(roll_mask_loss * self.roll_mask_loss_weight) if self.unmask_loss_weight > 0: # teacher head forward with torch.no_grad(): global_point_.feat = self.teacher.unmask_head(global_feat) # student forward local_point_ = self.student.backbone(local_point) local_point_ = self.up_cast(local_point_) local_point_ = self.fuse_representation(local_point_, self.student) unmask_pred_sim = self.student.unmask_head(local_point_.feat) with torch.no_grad(): principal_view_mask = global_point_.batch % self.num_global_view == 0 principal_view_batch = ( global_point_.batch[principal_view_mask] // self.num_global_view ) match_index = self.match_neighbour( local_point_.origin_coord, local_point_.offset[self.num_local_view - 1 :: self.num_local_view], global_point_.origin_coord[principal_view_mask], batch2offset(principal_view_batch), ) # teacher forward unmask_target_sim = self.sinkhorn_knopp( global_point_.feat[principal_view_mask][match_index[:, 1]], self.teacher_temp, ) # loss unmask_loss = -torch.sum( unmask_target_sim * F.log_softmax( unmask_pred_sim[match_index[:, 0]] / self.student_temp, dim=-1 ), dim=-1, ) unmask_loss = torch_scatter.segment_coo( unmask_loss, index=local_point_.batch[match_index[:, 0]], reduce="mean", ).mean() result_dict["unmask_loss"] = unmask_loss result_dict["loss"].append(unmask_loss * self.unmask_loss_weight) result_dict["loss"] = sum(result_dict["loss"]) result_dict['total_loss'] = result_dict['loss'].detach().clone() # sync component losses for logging if (ws:=get_world_size()) > 1: for key in list(result_dict.keys()): if key == 'loss': continue synced_loss = result_dict[key].detach() dist.all_reduce(synced_loss, op=dist.ReduceOp.SUM) synced_loss.div_(ws) result_dict[key] = synced_loss return result_dict