LibContinual / core /model /backbone /resnet_id.py
boringKey's picture
Sync from GitHub
3cdd985 verified
Raw
History Blame Contribute Delete
5.82 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
__all__ = ["resnet18_id2"]
def conv3x3(in_planes, out_planes, stride=1):
return nn.Conv2d(
in_planes,
out_planes,
kernel_size=3,
stride=stride,
padding=1,
bias=False,
)
class BasicBlock(nn.Module):
expansion = 1
def __init__(self, in_planes, planes, stride=1):
super().__init__()
self.conv1 = conv3x3(in_planes, planes, stride)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = conv3x3(planes, planes)
self.bn2 = nn.BatchNorm2d(planes)
self.shortcut = nn.Sequential()
if stride != 1 or in_planes != self.expansion * planes:
self.shortcut = nn.Sequential(
nn.Conv2d(
in_planes,
self.expansion * planes,
kernel_size=1,
stride=stride,
bias=False,
),
nn.BatchNorm2d(self.expansion * planes),
)
def forward(self, x):
out = F.relu(self.bn1(self.conv1(x)), inplace=True)
out = self.bn2(self.conv2(out))
out = out + self.shortcut(x)
return F.relu(out, inplace=True)
class CosineClassifier(nn.Module):
def __init__(self, feat_dim, num_classes, temperature=12.0):
super().__init__()
fc = nn.Linear(feat_dim, num_classes)
self.weight = nn.Parameter(fc.weight.t())
self.bias = nn.Parameter(fc.bias)
self.temperature = nn.Parameter(torch.tensor([temperature]), requires_grad=False)
def forward(self, features):
features = F.normalize(features, p=2, dim=1, eps=1e-12)
weight = F.normalize(self.weight, p=2, dim=0, eps=1e-12)
return self.temperature * torch.mm(features, weight)
class ResNetStage1(nn.Module):
def __init__(self, block, num_blocks, nf):
super().__init__()
self.in_planes = nf
self.conv1 = conv3x3(3, nf)
self.bn1 = nn.BatchNorm2d(nf)
self.layer1 = self._make_layer(block, nf, num_blocks[0], stride=1)
self.layer2 = self._make_layer(block, nf * 2, num_blocks[1], stride=2)
def _make_layer(self, block, planes, num_blocks, stride):
strides = [stride] + [1] * (num_blocks - 1)
layers = []
for stride in strides:
layers.append(block(self.in_planes, planes, stride))
self.in_planes = planes * block.expansion
return nn.Sequential(*layers)
def forward(self, x):
out = F.relu(self.bn1(self.conv1(x)), inplace=True)
out = self.layer1(out)
return self.layer2(out)
class ResNetStage2(nn.Module):
def __init__(self, block, num_blocks, num_classes, nf, use_cos=False):
super().__init__()
self.in_planes = nf * 2
self.num_classes = num_classes
self.out_dim = nf * 8 * block.expansion
self.label_fc = nn.Sequential(
nn.Linear(num_classes, nf * 2),
nn.LeakyReLU(inplace=True),
)
self.layer3 = self._make_layer(block, nf * 4, num_blocks[2], stride=2)
self.layer4 = self._make_layer(block, nf * 8, num_blocks[3], stride=2)
if use_cos:
self.classifier = CosineClassifier(self.out_dim, num_classes)
else:
self.classifier = nn.Linear(self.out_dim, num_classes)
def _make_layer(self, block, planes, num_blocks, stride):
strides = [stride] + [1] * (num_blocks - 1)
layers = []
for stride in strides:
layers.append(block(self.in_planes, planes, stride))
self.in_planes = planes * block.expansion
return nn.Sequential(*layers)
def forward(self, features, label_signal):
out = features + self.label_fc(label_signal)[..., None, None]
out = self.layer3(out)
out = self.layer4(out)
pooled = F.avg_pool2d(out, out.shape[2])
flat = torch.flatten(pooled, 1)
logits = self.classifier(flat)
return logits[:, : self.num_classes], flat, out
class IdempotentResNet(nn.Module):
"""ResNet-18 split for IDER.
The image path produces an intermediate feature map. The second input is a
label/probability signal projected to the same channel dimension and added
before the later ResNet stages, following the official IDER implementation.
"""
def __init__(self, block, num_blocks, num_classes=100, nf=64, use_cos=False, **kwargs):
super().__init__()
self.num_classes = num_classes
self.out_dim = nf * 8 * block.expansion
self.f1 = ResNetStage1(block, num_blocks, nf)
self.f2 = ResNetStage2(block, num_blocks, num_classes, nf, use_cos=use_cos)
def empty_signal(self, batch_size, device, upto=None):
upto = self.num_classes if upto is None else min(upto, self.num_classes)
signal = torch.zeros(batch_size, self.num_classes, device=device)
signal[:, :upto] = 1.0 / max(upto, 1)
return signal
def forward(self, x, y=None, upto=None, returnt="all"):
if y is None:
y = self.empty_signal(x.size(0), x.device, upto=upto)
stage1_features = self.f1(x)
logits, features, fmap = self.f2(stage1_features, y)
if returnt == "logits":
return logits
if returnt == "features":
return {"features": features}
return {"logits": logits, "features": features, "fmaps": [stage1_features, fmap]}
def resnet18_id2(pretrained=False, progress=True, num_classes=100, nf=64, use_cos=False, **kwargs):
if "nclasses" in kwargs:
num_classes = kwargs["nclasses"]
return IdempotentResNet(
BasicBlock,
[2, 2, 2, 2],
num_classes=num_classes,
nf=nf,
use_cos=use_cos,
)