"""
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