chore: Clean test directory (remove old folders/zips, rename db)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- test/{ab.nn.db.zst → ab.nn.zst-2.2.9} +0 -0
- test/lemur_data_sync.zip +0 -3
- test/nn/AirNet-626c3eb9-5c0a-43be-bcba-745592729769.py +0 -97
- test/nn/AirNet-777bc9dc-e6c7-4bff-8be1-a1b02adea68f.py +0 -97
- test/nn/AirNext-1c889567-e226-44b8-9ced-2cbb8ad0a561.py +0 -126
- test/nn/AirNext-31e78268-e095-43be-bc9b-ad8b34e76201.py +0 -128
- test/nn/AirNext-8c916d56-4362-4ab3-8f8a-94b73fe876fb.py +0 -126
- test/nn/AirNext.py +0 -125
- test/nn/AlexNet-69c52339-4eac-45f1-bbfe-2c51949701f1.py +0 -62
- test/nn/AlexNet-ad69700d-0e12-458f-afad-93f03988a4e7.py +0 -64
- test/nn/AlexNet-bb84fa5d-5dd8-4cc0-8a30-5c41958bcd94.py +0 -63
- test/nn/AlexNet.py +0 -61
- test/nn/BagNet-001bf3e2-17c2-4fdf-948e-493677e58a3b.py +0 -134
- test/nn/BagNet-560341b5-15a8-4829-a8ac-ea4c7391a950.py +0 -129
- test/nn/BagNet-6ecd3fc7-5ce2-4876-86a6-5e8250c43d78.py +0 -129
- test/nn/BagNet-7e541be1-6b60-445d-bbbf-3b655eeefc9a.py +0 -129
- test/nn/BagNet-7ebe6562-46c6-4406-96a2-bf3914ac8516.py +0 -139
- test/nn/BagNet-7f792262-31cf-477e-a78a-3494c122332d.py +0 -129
- test/nn/BayesianNet-024b0436-9ad5-4a1f-86d0-946e577ffc2d.py +0 -244
- test/nn/BayesianNet-0901ac22-d7f5-4deb-94c9-970e7955bd68.py +0 -242
- test/nn/BayesianNet-1.py +0 -241
- test/nn/BayesianNet-4f11c8da-cfe1-46ba-b5d0-b5d899929a2e.py +0 -238
- test/nn/C10C-RESNETLSTM-6a517327bf0ef897a22186a2061e85b3.py +0 -180
- test/nn/C10C-RESNETLSTM-8f7ac9c241d5b9546f8cd3484e0e100b.py +0 -245
- test/nn/C10C-RESNETLSTM-IMG-CAP-IMPROVED.py +0 -230
- test/nn/C10C-ResNetTransformer-187ccbee8050ac295637ecedecb4da1e.py +0 -193
- test/nn/C5C-RESNETLSTM-4.py +0 -222
- test/nn/C5C-RESNETLSTM-c42512d71480c8ef10f31e3e6c33bbdf.py +0 -150
- test/nn/C5C-ResNetTransformer-83fb6b6bb7c76b742ad0713d29463514.py +0 -181
- test/nn/C8C-ResNetTransformer-7730b6eb6979d27e2e1bbc7d05255dff.py +0 -239
- test/nn/ComplexNet.py +0 -295
- test/nn/ConditionalDiffusion.py +0 -230
- test/nn/ConditionalGAN.py +0 -278
- test/nn/ConditionalVAE3.py +0 -213
- test/nn/ConditionalVAE4.py +0 -268
- test/nn/ConvNeXt-dda5bf19-9ac1-460b-9bfd-735eec2f4904.py +0 -172
- test/nn/DPN107.py +0 -92
- test/nn/DPN131-8e6e495b-85cb-4a71-8b91-6d89372e0a0c.py +0 -86
- test/nn/DPN131-c53a40b8-b874-4c8b-999b-0944a1173a46.py +0 -86
- test/nn/DPN131-e8980802-6b89-4170-8608-327297706df0.py +0 -86
- test/nn/DPN131.py +0 -85
- test/nn/DPN68-9693aa0b-80bf-4393-a9e0-dd985a5ab128.py +0 -83
- test/nn/DPN68-c9cdb196-7596-4368-974a-56edf8b10381.py +0 -83
- test/nn/DarkNet-11e8caec-5e73-461e-a101-3aa39dfec644.py +0 -96
- test/nn/DarkNet-51277b91-c3a1-4669-9fb5-849ea97bd1b4.py +0 -95
- test/nn/DarkNet-d434ba1c-25ea-4160-a41d-4c477dba7bc0.py +0 -96
- test/nn/DarkNet.py +0 -95
- test/nn/DeepLabV3-1.py +0 -382
- test/nn/DeepLabV3-2.py +0 -382
- test/nn/DenoiseUNet.py +0 -135
test/{ab.nn.db.zst → ab.nn.zst-2.2.9}
RENAMED
|
File without changes
|
test/lemur_data_sync.zip
DELETED
|
@@ -1,3 +0,0 @@
|
|
| 1 |
-
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:59cb166870c79c9a1a75974977d97655016bcc088b042ea3ba339d560af7e7ba
|
| 3 |
-
size 229788994
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNet-626c3eb9-5c0a-43be-bcba-745592729769.py
DELETED
|
@@ -1,97 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class AirInitBlock(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels):
|
| 12 |
-
super().__init__()
|
| 13 |
-
self.layers = nn.Sequential(
|
| 14 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2, padding=1),
|
| 15 |
-
nn.BatchNorm2d(out_channels),
|
| 16 |
-
nn.ReLU(inplace=True)
|
| 17 |
-
)
|
| 18 |
-
|
| 19 |
-
def forward(self, x):
|
| 20 |
-
return self.layers(x)
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
class AirUnit(nn.Module):
|
| 24 |
-
def __init__(self, in_channels, out_channels, stride):
|
| 25 |
-
super().__init__()
|
| 26 |
-
self.layers = nn.Sequential(
|
| 27 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1),
|
| 28 |
-
nn.BatchNorm2d(out_channels),
|
| 29 |
-
nn.ReLU(inplace=True),
|
| 30 |
-
nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1),
|
| 31 |
-
nn.BatchNorm2d(out_channels)
|
| 32 |
-
)
|
| 33 |
-
self.downsample = (
|
| 34 |
-
nn.Sequential(
|
| 35 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=2, stride=stride, bias=False),
|
| 36 |
-
nn.BatchNorm2d(out_channels)
|
| 37 |
-
) if stride != 1 or in_channels != out_channels else nn.Identity()
|
| 38 |
-
)
|
| 39 |
-
self.relu = nn.ReLU(inplace=True)
|
| 40 |
-
|
| 41 |
-
def forward(self, x):
|
| 42 |
-
residual = self.downsample(x)
|
| 43 |
-
x = self.layers(x)
|
| 44 |
-
return self.relu(x + residual)
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
class Net(nn.Module):
|
| 48 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 49 |
-
super().__init__()
|
| 50 |
-
self.device = device
|
| 51 |
-
self.in_channels = in_shape[1]
|
| 52 |
-
self.image_size = in_shape[2]
|
| 53 |
-
self.num_classes = out_shape[0]
|
| 54 |
-
self.learning_rate = prm['lr']
|
| 55 |
-
self.momentum = prm['momentum']
|
| 56 |
-
|
| 57 |
-
channels = [64, 128, 256, 512]
|
| 58 |
-
init_block_channels = 64
|
| 59 |
-
|
| 60 |
-
self.features = self.build_features(init_block_channels, channels)
|
| 61 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 62 |
-
self.classifier = nn.Linear(channels[-1], self.num_classes)
|
| 63 |
-
|
| 64 |
-
def build_features(self, init_block_channels, channels):
|
| 65 |
-
layers = [AirInitBlock(self.in_channels, init_block_channels)]
|
| 66 |
-
for i, out_channels in enumerate(channels):
|
| 67 |
-
layers.append(AirUnit(
|
| 68 |
-
in_channels=init_block_channels if i == 0 else channels[i - 1],
|
| 69 |
-
out_channels=out_channels,
|
| 70 |
-
stride=1 if i == 0 else 2))
|
| 71 |
-
return nn.Sequential(*layers)
|
| 72 |
-
|
| 73 |
-
def forward(self, x):
|
| 74 |
-
x = self.features(x)
|
| 75 |
-
x = self.avgpool(x)
|
| 76 |
-
x = torch.flatten(x, 1)
|
| 77 |
-
return self.classifier(x)
|
| 78 |
-
|
| 79 |
-
def train_setup(self, prm):
|
| 80 |
-
self.to(self.device)
|
| 81 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 82 |
-
self.optimizer = torch.optim.SGD(
|
| 83 |
-
self.parameters(),
|
| 84 |
-
lr=self.learning_rate,
|
| 85 |
-
momentum=self.momentum
|
| 86 |
-
)
|
| 87 |
-
|
| 88 |
-
def learn(self, train_data):
|
| 89 |
-
self.train()
|
| 90 |
-
for inputs, labels in train_data:
|
| 91 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 92 |
-
self.optimizer.zero_grad()
|
| 93 |
-
outputs = self(inputs)
|
| 94 |
-
loss = self.criteria(outputs, labels)
|
| 95 |
-
loss.backward()
|
| 96 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 97 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNet-777bc9dc-e6c7-4bff-8be1-a1b02adea68f.py
DELETED
|
@@ -1,97 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class AirInitBlock(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels):
|
| 12 |
-
super().__init__()
|
| 13 |
-
self.layers = nn.Sequential(
|
| 14 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2, padding=1),
|
| 15 |
-
nn.BatchNorm2d(out_channels),
|
| 16 |
-
nn.ReLU(inplace=True)
|
| 17 |
-
)
|
| 18 |
-
|
| 19 |
-
def forward(self, x):
|
| 20 |
-
return self.layers(x)
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
class AirUnit(nn.Module):
|
| 24 |
-
def __init__(self, in_channels, out_channels, stride):
|
| 25 |
-
super().__init__()
|
| 26 |
-
self.layers = nn.Sequential(
|
| 27 |
-
nn.Conv2d(in_channels, 3, kernel_size=3, stride=stride, padding=1),
|
| 28 |
-
nn.BatchNorm2d(3),
|
| 29 |
-
nn.ReLU(inplace=True),
|
| 30 |
-
nn.Conv2d(3, out_channels, kernel_size=3, stride=1, padding=1),
|
| 31 |
-
nn.BatchNorm2d(out_channels)
|
| 32 |
-
)
|
| 33 |
-
self.downsample = (
|
| 34 |
-
nn.Sequential(
|
| 35 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),
|
| 36 |
-
nn.BatchNorm2d(out_channels)
|
| 37 |
-
) if stride != 1 or in_channels != out_channels else nn.Identity()
|
| 38 |
-
)
|
| 39 |
-
self.relu = nn.ReLU(inplace=True)
|
| 40 |
-
|
| 41 |
-
def forward(self, x):
|
| 42 |
-
residual = self.downsample(x)
|
| 43 |
-
x = self.layers(x)
|
| 44 |
-
return self.relu(x + residual)
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
class Net(nn.Module):
|
| 48 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 49 |
-
super().__init__()
|
| 50 |
-
self.device = device
|
| 51 |
-
self.in_channels = in_shape[1]
|
| 52 |
-
self.image_size = in_shape[2]
|
| 53 |
-
self.num_classes = out_shape[0]
|
| 54 |
-
self.learning_rate = prm['lr']
|
| 55 |
-
self.momentum = prm['momentum']
|
| 56 |
-
|
| 57 |
-
channels = [64, 128, 256, 512]
|
| 58 |
-
init_block_channels = 64
|
| 59 |
-
|
| 60 |
-
self.features = self.build_features(init_block_channels, channels)
|
| 61 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 62 |
-
self.classifier = nn.Linear(channels[-1], self.num_classes)
|
| 63 |
-
|
| 64 |
-
def build_features(self, init_block_channels, channels):
|
| 65 |
-
layers = [AirInitBlock(self.in_channels, init_block_channels)]
|
| 66 |
-
for i, out_channels in enumerate(channels):
|
| 67 |
-
layers.append(AirUnit(
|
| 68 |
-
in_channels=init_block_channels if i == 0 else channels[i - 1],
|
| 69 |
-
out_channels=out_channels,
|
| 70 |
-
stride=1 if i == 0 else 2))
|
| 71 |
-
return nn.Sequential(*layers)
|
| 72 |
-
|
| 73 |
-
def forward(self, x):
|
| 74 |
-
x = self.features(x)
|
| 75 |
-
x = self.avgpool(x)
|
| 76 |
-
x = torch.flatten(x, 1)
|
| 77 |
-
return self.classifier(x)
|
| 78 |
-
|
| 79 |
-
def train_setup(self, prm):
|
| 80 |
-
self.to(self.device)
|
| 81 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 82 |
-
self.optimizer = torch.optim.SGD(
|
| 83 |
-
self.parameters(),
|
| 84 |
-
lr=self.learning_rate,
|
| 85 |
-
momentum=self.momentum
|
| 86 |
-
)
|
| 87 |
-
|
| 88 |
-
def learn(self, train_data):
|
| 89 |
-
self.train()
|
| 90 |
-
for inputs, labels in train_data:
|
| 91 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 92 |
-
self.optimizer.zero_grad()
|
| 93 |
-
outputs = self(inputs)
|
| 94 |
-
loss = self.criteria(outputs, labels)
|
| 95 |
-
loss.backward()
|
| 96 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 97 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNext-1c889567-e226-44b8-9ced-2cbb8ad0a561.py
DELETED
|
@@ -1,126 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
import math
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
class AirBlock(nn.Module):
|
| 9 |
-
def __init__(self, in_channels, out_channels, groups=2, ratio=3):
|
| 10 |
-
super(AirBlock, self).__init__()
|
| 11 |
-
mid_channels = out_channels // ratio
|
| 12 |
-
self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 13 |
-
self.bn1 = nn.BatchNorm2d(mid_channels)
|
| 14 |
-
self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 15 |
-
self.conv2 = nn.Conv2d(mid_channels, mid_channels, kernel_size=3, stride=1, padding=1, groups=groups, bias=False)
|
| 16 |
-
self.bn2 = nn.BatchNorm2d(mid_channels)
|
| 17 |
-
self.conv3 = nn.Conv2d(mid_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 18 |
-
self.bn3 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
self.sigmoid = nn.Sigmoid()
|
| 20 |
-
|
| 21 |
-
def forward(self, x):
|
| 22 |
-
x = torch.relu(self.bn1(self.conv1(x)))
|
| 23 |
-
x = self.pool(x)
|
| 24 |
-
x = torch.relu(self.bn2(self.conv2(x)))
|
| 25 |
-
x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=True)
|
| 26 |
-
x = self.bn3(self.conv3(x))
|
| 27 |
-
x = self.sigmoid(x)
|
| 28 |
-
return x
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class AirNeXtUnit(nn.Module):
|
| 32 |
-
def __init__(self, in_channels, out_channels, stride, cardinality, bottleneck_width, ratio):
|
| 33 |
-
super(AirNeXtUnit, self).__init__()
|
| 34 |
-
mid_channels = out_channels // 4
|
| 35 |
-
D = int(math.floor(mid_channels * (bottleneck_width / 64.0)))
|
| 36 |
-
group_width = cardinality * D
|
| 37 |
-
self.use_air_block = (stride == 1 and mid_channels < 512)
|
| 38 |
-
|
| 39 |
-
self.conv1 = nn.Conv2d(in_channels, group_width, kernel_size=1, stride=1, padding=0, bias=False)
|
| 40 |
-
self.conv2 = nn.Conv2d(group_width, group_width, kernel_size=3, stride=stride, padding=1, groups=cardinality, bias=False)
|
| 41 |
-
self.conv3 = nn.Conv2d(group_width, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 42 |
-
if self.use_air_block:
|
| 43 |
-
self.air = AirBlock(in_channels, group_width, groups=cardinality // ratio, ratio=ratio)
|
| 44 |
-
|
| 45 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def forward(self, x):
|
| 51 |
-
if self.use_air_block:
|
| 52 |
-
att = self.air(x)
|
| 53 |
-
att = F.interpolate(att, size=x.shape[2:], mode="bilinear", align_corners=True) # Ensure att matches x dimensions
|
| 54 |
-
identity = self.identity_conv(x) if self.resize_identity else x
|
| 55 |
-
x = self.conv1(x)
|
| 56 |
-
x = self.conv2(x)
|
| 57 |
-
if self.use_air_block:
|
| 58 |
-
x = x * att
|
| 59 |
-
x = self.conv3(x)
|
| 60 |
-
x = x + identity
|
| 61 |
-
x = self.activ(x)
|
| 62 |
-
return x
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
class Net(nn.Module):
|
| 66 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 67 |
-
super(Net, self).__init__()
|
| 68 |
-
self.device = device
|
| 69 |
-
channel_number = in_shape[1]
|
| 70 |
-
image_size = in_shape[2]
|
| 71 |
-
class_number = out_shape[0]
|
| 72 |
-
|
| 73 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 74 |
-
init_block_channels = 64
|
| 75 |
-
cardinality = 32
|
| 76 |
-
bottleneck_width = 4
|
| 77 |
-
ratio = 2
|
| 78 |
-
|
| 79 |
-
self.in_size = image_size
|
| 80 |
-
self.num_classes = class_number
|
| 81 |
-
|
| 82 |
-
self.features = nn.Sequential(
|
| 83 |
-
nn.Conv2d(channel_number, init_block_channels, kernel_size=7, stride=2, padding=3, bias=False),
|
| 84 |
-
nn.BatchNorm2d(init_block_channels),
|
| 85 |
-
nn.ReLU(inplace=True),
|
| 86 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 87 |
-
)
|
| 88 |
-
in_channels = init_block_channels
|
| 89 |
-
for i, channels_per_stage in enumerate(channels):
|
| 90 |
-
stage = nn.Sequential()
|
| 91 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 92 |
-
stride = 2 if (j == 0) and (i != 0) else 1
|
| 93 |
-
stage.add_module("unit{}".format(j + 1), AirNeXtUnit(in_channels, out_channels, stride, cardinality, bottleneck_width, ratio))
|
| 94 |
-
in_channels = out_channels
|
| 95 |
-
self.features.add_module("stage{}".format(i + 1), stage)
|
| 96 |
-
|
| 97 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 98 |
-
self.output = nn.Linear(in_channels, class_number)
|
| 99 |
-
|
| 100 |
-
def forward(self, x):
|
| 101 |
-
x = self.features(x)
|
| 102 |
-
x = x.view(x.size(0), -1)
|
| 103 |
-
x = self.output(x)
|
| 104 |
-
return x
|
| 105 |
-
|
| 106 |
-
def train_setup(self, prm):
|
| 107 |
-
self.to(self.device)
|
| 108 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 109 |
-
self.optimizer = torch.optim.Adam(self.parameters(), lr=prm['lr'], weight_decay=1e-4)
|
| 110 |
-
self.scheduler = torch.optim.lr_scheduler.StepLR(self.optimizer, step_size=5, gamma=0.5)
|
| 111 |
-
|
| 112 |
-
def learn(self, train_data):
|
| 113 |
-
self.train()
|
| 114 |
-
for inputs, labels in train_data:
|
| 115 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 116 |
-
self.optimizer.zero_grad()
|
| 117 |
-
outputs = self(inputs)
|
| 118 |
-
loss = self.criteria(outputs, labels)
|
| 119 |
-
loss.backward()
|
| 120 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 121 |
-
self.optimizer.step()
|
| 122 |
-
self.scheduler.step()
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
def supported_hyperparameters():
|
| 126 |
-
return {'lr', 'momentum', 'dropout'}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNext-31e78268-e095-43be-bc9b-ad8b34e76201.py
DELETED
|
@@ -1,128 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
import math
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
class AirBlock(nn.Module):
|
| 9 |
-
def __init__(self, in_channels, out_channels, groups=1, ratio=2):
|
| 10 |
-
super(AirBlock, self).__init__()
|
| 11 |
-
mid_channels = out_channels // ratio
|
| 12 |
-
self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 13 |
-
self.bn1 = nn.BatchNorm2d(mid_channels)
|
| 14 |
-
self.pool = nn.MaxPool2d(kernel_size=3, stride=1, padding=1)
|
| 15 |
-
self.conv2 = nn.Conv2d(mid_channels, mid_channels, kernel_size=3, stride=1, padding=1, groups=groups, bias=False)
|
| 16 |
-
self.bn2 = nn.BatchNorm2d(mid_channels)
|
| 17 |
-
self.conv3 = nn.Conv2d(mid_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 18 |
-
self.bn3 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
self.sigmoid = nn.Sigmoid()
|
| 20 |
-
|
| 21 |
-
def forward(self, x):
|
| 22 |
-
x = torch.relu(self.bn1(self.conv1(x)))
|
| 23 |
-
x = self.pool(x)
|
| 24 |
-
x = torch.relu(self.bn2(self.conv2(x)))
|
| 25 |
-
x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=True)
|
| 26 |
-
x = self.bn3(self.conv3(x))
|
| 27 |
-
x = self.sigmoid(x)
|
| 28 |
-
return x
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class AirNeXtUnit(nn.Module):
|
| 32 |
-
def __init__(self, in_channels, out_channels, stride, cardinality, bottleneck_width, ratio):
|
| 33 |
-
super(AirNeXtUnit, self).__init__()
|
| 34 |
-
mid_channels = out_channels // 4
|
| 35 |
-
D = int(math.floor(mid_channels * (bottleneck_width / 64.0)))
|
| 36 |
-
group_width = cardinality * D
|
| 37 |
-
self.use_air_block = (stride == 1 and mid_channels < 512)
|
| 38 |
-
|
| 39 |
-
self.conv1 = nn.Conv2d(in_channels, group_width, kernel_size=1, stride=1, padding=0, bias=False)
|
| 40 |
-
self.conv2 = nn.Conv2d(group_width, group_width, kernel_size=3, stride=stride, padding=1, groups=cardinality, bias=False)
|
| 41 |
-
self.conv3 = nn.Conv2d(group_width, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 42 |
-
if self.use_air_block:
|
| 43 |
-
self.air = AirBlock(in_channels, group_width, groups=(cardinality // ratio), ratio=ratio)
|
| 44 |
-
|
| 45 |
-
self.resize_identity = (in_channels!= out_channels) or (stride!= 1)
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def forward(self, x):
|
| 51 |
-
if self.use_air_block:
|
| 52 |
-
att = self.air(x)
|
| 53 |
-
att = F.interpolate(att, size=x.shape[2:], mode="bilinear", align_corners=True) # Ensure att matches x dimensions
|
| 54 |
-
identity = self.identity_conv(x) if self.resize_identity else x
|
| 55 |
-
x = self.conv1(x)
|
| 56 |
-
x = self.conv2(x)
|
| 57 |
-
if self.use_air_block:
|
| 58 |
-
x = x * att
|
| 59 |
-
x = self.conv3(x)
|
| 60 |
-
x = x + identity
|
| 61 |
-
x = self.activ(x)
|
| 62 |
-
return x
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
class Net(nn.Module):
|
| 66 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 67 |
-
super(Net, self).__init__()
|
| 68 |
-
self.device = device
|
| 69 |
-
channel_number = in_shape[1]
|
| 70 |
-
image_size = in_shape[2]
|
| 71 |
-
class_number = out_shape[0]
|
| 72 |
-
|
| 73 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256, 512], [512, 512, 512, 512]]
|
| 74 |
-
init_block_channels = 64
|
| 75 |
-
cardinality = 32
|
| 76 |
-
bottleneck_width = 4
|
| 77 |
-
ratio = 2
|
| 78 |
-
|
| 79 |
-
self.in_size = image_size
|
| 80 |
-
self.num_classes = class_number
|
| 81 |
-
|
| 82 |
-
self.features = nn.Sequential(
|
| 83 |
-
nn.Conv2d(channel_number, init_block_channels, kernel_size=7, stride=2, padding=3, bias=False),
|
| 84 |
-
nn.BatchNorm2d(init_block_channels),
|
| 85 |
-
nn.ReLU(inplace=True),
|
| 86 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 87 |
-
)
|
| 88 |
-
in_channels = init_block_channels
|
| 89 |
-
self.in_channels = in_channels
|
| 90 |
-
for i, channels_per_stage in enumerate(channels):
|
| 91 |
-
stage = nn.Sequential()
|
| 92 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 93 |
-
stride = 2 if (j == 0) and (i!= 0) else 1
|
| 94 |
-
stage.add_module("unit{}".format(j + 1), AirNeXtUnit(in_channels, out_channels, stride, cardinality, bottleneck_width, ratio))
|
| 95 |
-
in_channels = out_channels
|
| 96 |
-
self.features.add_module("stage{}".format(i + 1), stage)
|
| 97 |
-
|
| 98 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 99 |
-
self.output = nn.Linear(in_channels, class_number)
|
| 100 |
-
|
| 101 |
-
def forward(self, x):
|
| 102 |
-
x = self.features(x)
|
| 103 |
-
x = x.view(x.size(0), -1)
|
| 104 |
-
x = self.output(x)
|
| 105 |
-
return x
|
| 106 |
-
|
| 107 |
-
def train_setup(self, prm):
|
| 108 |
-
self.to(self.device)
|
| 109 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 110 |
-
self.optimizer = torch.optim.Adam(self.parameters(), lr=prm['lr'], weight_decay=1e-4)
|
| 111 |
-
self.scheduler = torch.optim.lr_scheduler.StepLR(self.optimizer, step_size=5, gamma=0.5)
|
| 112 |
-
|
| 113 |
-
def learn(self, train_data):
|
| 114 |
-
self.train()
|
| 115 |
-
for inputs, labels in train_data:
|
| 116 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 117 |
-
self.optimizer.zero_grad()
|
| 118 |
-
outputs = self(inputs)
|
| 119 |
-
loss = self.criteria(outputs, labels)
|
| 120 |
-
loss.backward()
|
| 121 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 122 |
-
self.optimizer.step()
|
| 123 |
-
self.scheduler.step()
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
def supported_hyperparameters():
|
| 127 |
-
return {'lr','momentum', 'dropout'}
|
| 128 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNext-8c916d56-4362-4ab3-8f8a-94b73fe876fb.py
DELETED
|
@@ -1,126 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
import math
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
class AirBlock(nn.Module):
|
| 9 |
-
def __init__(self, in_channels, out_channels, groups=3, ratio=4):
|
| 10 |
-
super(AirBlock, self).__init__()
|
| 11 |
-
mid_channels = out_channels // ratio
|
| 12 |
-
self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 13 |
-
self.bn1 = nn.BatchNorm2d(mid_channels)
|
| 14 |
-
self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 15 |
-
self.conv2 = nn.Conv2d(mid_channels, mid_channels, kernel_size=3, stride=1, padding=1, groups=groups, bias=False)
|
| 16 |
-
self.bn2 = nn.BatchNorm2d(mid_channels)
|
| 17 |
-
self.conv3 = nn.Conv2d(mid_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 18 |
-
self.bn3 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
self.sigmoid = nn.Sigmoid()
|
| 20 |
-
|
| 21 |
-
def forward(self, x):
|
| 22 |
-
x = torch.relu(self.bn1(self.conv1(x)))
|
| 23 |
-
x = self.pool(x)
|
| 24 |
-
x = torch.relu(self.bn2(self.conv2(x)))
|
| 25 |
-
x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=True)
|
| 26 |
-
x = self.bn3(self.conv3(x))
|
| 27 |
-
x = self.sigmoid(x)
|
| 28 |
-
return x
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class AirNeXtUnit(nn.Module):
|
| 32 |
-
def __init__(self, in_channels, out_channels, stride, cardinality=8, bottleneck_width=4, ratio=3):
|
| 33 |
-
super(AirNeXtUnit, self).__init__()
|
| 34 |
-
mid_channels = out_channels // 4
|
| 35 |
-
D = int(math.floor(mid_channels * (bottleneck_width / 64.0)))
|
| 36 |
-
group_width = cardinality * D
|
| 37 |
-
self.use_air_block = (stride == 1 and mid_channels < 512)
|
| 38 |
-
|
| 39 |
-
self.conv1 = nn.Conv2d(in_channels, group_width, kernel_size=1, stride=1, padding=0, bias=False)
|
| 40 |
-
self.conv2 = nn.Conv2d(group_width, group_width, kernel_size=3, stride=stride, padding=1, groups=cardinality, bias=False)
|
| 41 |
-
self.conv3 = nn.Conv2d(group_width, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 42 |
-
if self.use_air_block:
|
| 43 |
-
self.air = AirBlock(in_channels, group_width, groups=(cardinality // ratio), ratio=ratio)
|
| 44 |
-
|
| 45 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def forward(self, x):
|
| 51 |
-
if self.use_air_block:
|
| 52 |
-
att = self.air(x)
|
| 53 |
-
att = F.interpolate(att, size=x.shape[2:], mode="bilinear", align_corners=True) # Ensure att matches x dimensions
|
| 54 |
-
identity = self.identity_conv(x) if self.resize_identity else x
|
| 55 |
-
x = self.conv1(x)
|
| 56 |
-
x = self.conv2(x)
|
| 57 |
-
if self.use_air_block:
|
| 58 |
-
x = x * att
|
| 59 |
-
x = self.conv3(x)
|
| 60 |
-
x = x + identity
|
| 61 |
-
x = self.activ(x)
|
| 62 |
-
return x
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
class Net(nn.Module):
|
| 66 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 67 |
-
super(Net, self).__init__()
|
| 68 |
-
self.device = device
|
| 69 |
-
channel_number = in_shape[1]
|
| 70 |
-
image_size = in_shape[2]
|
| 71 |
-
class_number = out_shape[0]
|
| 72 |
-
|
| 73 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 74 |
-
init_block_channels = 64
|
| 75 |
-
cardinality = 32
|
| 76 |
-
bottleneck_width = 4
|
| 77 |
-
ratio = 2
|
| 78 |
-
|
| 79 |
-
self.in_size = image_size
|
| 80 |
-
self.num_classes = class_number
|
| 81 |
-
|
| 82 |
-
self.features = nn.Sequential(
|
| 83 |
-
nn.Conv2d(channel_number, init_block_channels, kernel_size=7, stride=2, padding=3, bias=False),
|
| 84 |
-
nn.BatchNorm2d(init_block_channels),
|
| 85 |
-
nn.ReLU(inplace=True),
|
| 86 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 87 |
-
)
|
| 88 |
-
in_channels = init_block_channels
|
| 89 |
-
for i, channels_per_stage in enumerate(channels):
|
| 90 |
-
stage = nn.Sequential()
|
| 91 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 92 |
-
stride = 2 if (j == 0) and (i != 0) else 1
|
| 93 |
-
stage.add_module("unit{}".format(j + 1), AirNeXtUnit(in_channels, out_channels, stride, cardinality, bottleneck_width, ratio))
|
| 94 |
-
in_channels = out_channels
|
| 95 |
-
self.features.add_module("stage{}".format(i + 1), stage)
|
| 96 |
-
|
| 97 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 98 |
-
self.output = nn.Linear(in_channels, class_number)
|
| 99 |
-
|
| 100 |
-
def forward(self, x):
|
| 101 |
-
x = self.features(x)
|
| 102 |
-
x = x.view(x.size(0), -1)
|
| 103 |
-
x = self.output(x)
|
| 104 |
-
return x
|
| 105 |
-
|
| 106 |
-
def train_setup(self, prm):
|
| 107 |
-
self.to(self.device)
|
| 108 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 109 |
-
self.optimizer = torch.optim.Adam(self.parameters(), lr=prm['lr'], weight_decay=1e-4)
|
| 110 |
-
self.scheduler = torch.optim.lr_scheduler.StepLR(self.optimizer, step_size=5, gamma=0.5)
|
| 111 |
-
|
| 112 |
-
def learn(self, train_data):
|
| 113 |
-
self.train()
|
| 114 |
-
for inputs, labels in train_data:
|
| 115 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 116 |
-
self.optimizer.zero_grad()
|
| 117 |
-
outputs = self(inputs)
|
| 118 |
-
loss = self.criteria(outputs, labels)
|
| 119 |
-
loss.backward()
|
| 120 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 121 |
-
self.optimizer.step()
|
| 122 |
-
self.scheduler.step()
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
def supported_hyperparameters():
|
| 126 |
-
return {'lr', 'momentum', 'dropout'}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AirNext.py
DELETED
|
@@ -1,125 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
import math
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
class AirBlock(nn.Module):
|
| 8 |
-
def __init__(self, in_channels, out_channels, groups=1, ratio=2):
|
| 9 |
-
super(AirBlock, self).__init__()
|
| 10 |
-
mid_channels = out_channels // ratio
|
| 11 |
-
self.conv1 = nn.Conv2d(in_channels, mid_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 12 |
-
self.bn1 = nn.BatchNorm2d(mid_channels)
|
| 13 |
-
self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 14 |
-
self.conv2 = nn.Conv2d(mid_channels, mid_channels, kernel_size=3, stride=1, padding=1, groups=groups, bias=False)
|
| 15 |
-
self.bn2 = nn.BatchNorm2d(mid_channels)
|
| 16 |
-
self.conv3 = nn.Conv2d(mid_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 17 |
-
self.bn3 = nn.BatchNorm2d(out_channels)
|
| 18 |
-
self.sigmoid = nn.Sigmoid()
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
x = torch.relu(self.bn1(self.conv1(x)))
|
| 22 |
-
x = self.pool(x)
|
| 23 |
-
x = torch.relu(self.bn2(self.conv2(x)))
|
| 24 |
-
x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=True)
|
| 25 |
-
x = self.bn3(self.conv3(x))
|
| 26 |
-
x = self.sigmoid(x)
|
| 27 |
-
return x
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
class AirNeXtUnit(nn.Module):
|
| 31 |
-
def __init__(self, in_channels, out_channels, stride, cardinality, bottleneck_width, ratio):
|
| 32 |
-
super(AirNeXtUnit, self).__init__()
|
| 33 |
-
mid_channels = out_channels // 4
|
| 34 |
-
D = int(math.floor(mid_channels * (bottleneck_width / 64.0)))
|
| 35 |
-
group_width = cardinality * D
|
| 36 |
-
self.use_air_block = (stride == 1 and mid_channels < 512)
|
| 37 |
-
|
| 38 |
-
self.conv1 = nn.Conv2d(in_channels, group_width, kernel_size=1, stride=1, padding=0, bias=False)
|
| 39 |
-
self.conv2 = nn.Conv2d(group_width, group_width, kernel_size=3, stride=stride, padding=1, groups=cardinality, bias=False)
|
| 40 |
-
self.conv3 = nn.Conv2d(group_width, out_channels, kernel_size=1, stride=1, padding=0, bias=False)
|
| 41 |
-
if self.use_air_block:
|
| 42 |
-
self.air = AirBlock(in_channels, group_width, groups=(cardinality // ratio), ratio=ratio)
|
| 43 |
-
|
| 44 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 45 |
-
if self.resize_identity:
|
| 46 |
-
self.identity_conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False)
|
| 47 |
-
self.activ = nn.ReLU(inplace=True)
|
| 48 |
-
|
| 49 |
-
def forward(self, x):
|
| 50 |
-
if self.use_air_block:
|
| 51 |
-
att = self.air(x)
|
| 52 |
-
att = F.interpolate(att, size=x.shape[2:], mode="bilinear", align_corners=True) # Ensure att matches x dimensions
|
| 53 |
-
identity = self.identity_conv(x) if self.resize_identity else x
|
| 54 |
-
x = self.conv1(x)
|
| 55 |
-
x = self.conv2(x)
|
| 56 |
-
if self.use_air_block:
|
| 57 |
-
x = x * att
|
| 58 |
-
x = self.conv3(x)
|
| 59 |
-
x = x + identity
|
| 60 |
-
x = self.activ(x)
|
| 61 |
-
return x
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
class Net(nn.Module):
|
| 65 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 66 |
-
super(Net, self).__init__()
|
| 67 |
-
self.device = device
|
| 68 |
-
channel_number = in_shape[1]
|
| 69 |
-
image_size = in_shape[2]
|
| 70 |
-
class_number = out_shape[0]
|
| 71 |
-
|
| 72 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 73 |
-
init_block_channels = 64
|
| 74 |
-
cardinality = 32
|
| 75 |
-
bottleneck_width = 4
|
| 76 |
-
ratio = 2
|
| 77 |
-
|
| 78 |
-
self.in_size = image_size
|
| 79 |
-
self.num_classes = class_number
|
| 80 |
-
|
| 81 |
-
self.features = nn.Sequential(
|
| 82 |
-
nn.Conv2d(channel_number, init_block_channels, kernel_size=7, stride=2, padding=3, bias=False),
|
| 83 |
-
nn.BatchNorm2d(init_block_channels),
|
| 84 |
-
nn.ReLU(inplace=True),
|
| 85 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 86 |
-
)
|
| 87 |
-
in_channels = init_block_channels
|
| 88 |
-
for i, channels_per_stage in enumerate(channels):
|
| 89 |
-
stage = nn.Sequential()
|
| 90 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 91 |
-
stride = 2 if (j == 0) and (i != 0) else 1
|
| 92 |
-
stage.add_module("unit{}".format(j + 1), AirNeXtUnit(in_channels, out_channels, stride, cardinality, bottleneck_width, ratio))
|
| 93 |
-
in_channels = out_channels
|
| 94 |
-
self.features.add_module("stage{}".format(i + 1), stage)
|
| 95 |
-
|
| 96 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 97 |
-
self.output = nn.Linear(in_channels, class_number)
|
| 98 |
-
|
| 99 |
-
def forward(self, x):
|
| 100 |
-
x = self.features(x)
|
| 101 |
-
x = x.view(x.size(0), -1)
|
| 102 |
-
x = self.output(x)
|
| 103 |
-
return x
|
| 104 |
-
|
| 105 |
-
def train_setup(self, prm):
|
| 106 |
-
self.to(self.device)
|
| 107 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 108 |
-
self.optimizer = torch.optim.Adam(self.parameters(), lr=prm['lr'], weight_decay=1e-4)
|
| 109 |
-
self.scheduler = torch.optim.lr_scheduler.StepLR(self.optimizer, step_size=5, gamma=0.5)
|
| 110 |
-
|
| 111 |
-
def learn(self, train_data):
|
| 112 |
-
self.train()
|
| 113 |
-
for inputs, labels in train_data:
|
| 114 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 115 |
-
self.optimizer.zero_grad()
|
| 116 |
-
outputs = self(inputs)
|
| 117 |
-
loss = self.criteria(outputs, labels)
|
| 118 |
-
loss.backward()
|
| 119 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 120 |
-
self.optimizer.step()
|
| 121 |
-
self.scheduler.step()
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
def supported_hyperparameters():
|
| 125 |
-
return {'lr', 'momentum', 'dropout'}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AlexNet-69c52339-4eac-45f1-bbfe-2c51949701f1.py
DELETED
|
@@ -1,62 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class Net(nn.Module):
|
| 11 |
-
|
| 12 |
-
def train_setup(self, prm):
|
| 13 |
-
self.to(self.device)
|
| 14 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 15 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 16 |
-
|
| 17 |
-
def learn(self, train_data):
|
| 18 |
-
for inputs, labels in train_data:
|
| 19 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 20 |
-
self.optimizer.zero_grad()
|
| 21 |
-
outputs = self(inputs)
|
| 22 |
-
loss = self.criteria[0](outputs, labels)
|
| 23 |
-
loss.backward()
|
| 24 |
-
nn.utils.clip_grad_norm_(self.parameters(), 2) # Changed from 3 to 2
|
| 25 |
-
self.optimizer.step()
|
| 26 |
-
|
| 27 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 28 |
-
super().__init__()
|
| 29 |
-
self.device = device
|
| 30 |
-
self.features = nn.Sequential(
|
| 31 |
-
nn.Conv2d(in_shape[1], 64, kernel_size=7, stride=4, padding=2), # Changed from 11 to 7
|
| 32 |
-
nn.ReLU(inplace=True),
|
| 33 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 34 |
-
nn.Conv2d(64, 192, kernel_size=5, padding=2),
|
| 35 |
-
nn.ReLU(inplace=True),
|
| 36 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 37 |
-
nn.Conv2d(192, 384, kernel_size=3, padding=1),
|
| 38 |
-
nn.ReLU(inplace=True),
|
| 39 |
-
nn.Conv2d(384, 256, kernel_size=3, padding=1),
|
| 40 |
-
nn.ReLU(inplace=True),
|
| 41 |
-
nn.Conv2d(256, 256, kernel_size=3, padding=1),
|
| 42 |
-
nn.ReLU(inplace=True),
|
| 43 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 44 |
-
)
|
| 45 |
-
dropout: float = prm['dropout']
|
| 46 |
-
self.avgpool = nn.AdaptiveAvgPool2d((6, 6))
|
| 47 |
-
self.classifier = nn.Sequential(
|
| 48 |
-
nn.Dropout(p=dropout),
|
| 49 |
-
nn.Linear(256 * 6 * 6, 4096),
|
| 50 |
-
nn.ReLU(inplace=True),
|
| 51 |
-
nn.Dropout(p=dropout),
|
| 52 |
-
nn.Linear(4096, 4096),
|
| 53 |
-
nn.ReLU(inplace=True),
|
| 54 |
-
nn.Linear(4096, out_shape[0]),
|
| 55 |
-
)
|
| 56 |
-
|
| 57 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 58 |
-
x = self.features(x)
|
| 59 |
-
x = self.avgpool(x)
|
| 60 |
-
x = torch.flatten(x, 1)
|
| 61 |
-
x = self.classifier(x)
|
| 62 |
-
return x
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AlexNet-ad69700d-0e12-458f-afad-93f03988a4e7.py
DELETED
|
@@ -1,64 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class Net(nn.Module):
|
| 11 |
-
|
| 12 |
-
def train_setup(self, prm):
|
| 13 |
-
self.to(self.device)
|
| 14 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 15 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 16 |
-
|
| 17 |
-
def learn(self, train_data):
|
| 18 |
-
for inputs, labels in train_data:
|
| 19 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 20 |
-
self.optimizer.zero_grad()
|
| 21 |
-
outputs = self(inputs)
|
| 22 |
-
loss = self.criteria[0](outputs, labels)
|
| 23 |
-
loss.backward()
|
| 24 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 25 |
-
self.optimizer.step()
|
| 26 |
-
|
| 27 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 28 |
-
super().__init__()
|
| 29 |
-
self.device = device
|
| 30 |
-
self.features = nn.Sequential(
|
| 31 |
-
nn.Conv2d(in_shape[1], 64, kernel_size=11, stride=4, padding=2),
|
| 32 |
-
nn.ReLU(inplace=True),
|
| 33 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 34 |
-
nn.Conv2d(64, 256, kernel_size=5, padding=2), # Changed from 192 to 256
|
| 35 |
-
nn.ReLU(inplace=True),
|
| 36 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 37 |
-
nn.Conv2d(256, 384, kernel_size=3, padding=1),
|
| 38 |
-
nn.ReLU(inplace=True),
|
| 39 |
-
nn.Conv2d(384, 256, kernel_size=3, padding=1), # Changed from 384 to 256
|
| 40 |
-
nn.ReLU(inplace=True),
|
| 41 |
-
nn.Conv2d(256, 256, kernel_size=3, padding=1),
|
| 42 |
-
nn.ReLU(inplace=True),
|
| 43 |
-
nn.Conv2d(256, 256, kernel_size=3, padding=1),
|
| 44 |
-
nn.ReLU(inplace=True),
|
| 45 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 46 |
-
)
|
| 47 |
-
dropout: float = prm['dropout']
|
| 48 |
-
self.avgpool = nn.AdaptiveAvgPool2d((6, 6))
|
| 49 |
-
self.classifier = nn.Sequential(
|
| 50 |
-
nn.Dropout(p=dropout),
|
| 51 |
-
nn.Linear(256 * 6 * 6, 4096),
|
| 52 |
-
nn.ReLU(inplace=True),
|
| 53 |
-
nn.Dropout(p=dropout),
|
| 54 |
-
nn.Linear(4096, 4096),
|
| 55 |
-
nn.ReLU(inplace=True),
|
| 56 |
-
nn.Linear(4096, out_shape[0]),
|
| 57 |
-
)
|
| 58 |
-
|
| 59 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 60 |
-
x = self.features(x)
|
| 61 |
-
x = self.avgpool(x)
|
| 62 |
-
x = torch.flatten(x, 1)
|
| 63 |
-
x = self.classifier(x)
|
| 64 |
-
return x
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AlexNet-bb84fa5d-5dd8-4cc0-8a30-5c41958bcd94.py
DELETED
|
@@ -1,63 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class Net(nn.Module):
|
| 11 |
-
|
| 12 |
-
def train_setup(self, prm):
|
| 13 |
-
self.to(self.device)
|
| 14 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 15 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 16 |
-
|
| 17 |
-
def learn(self, train_data):
|
| 18 |
-
for inputs, labels in train_data:
|
| 19 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 20 |
-
self.optimizer.zero_grad()
|
| 21 |
-
outputs = self(inputs)
|
| 22 |
-
loss = self.criteria[0](outputs, labels)
|
| 23 |
-
loss.backward()
|
| 24 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 25 |
-
self.optimizer.step()
|
| 26 |
-
|
| 27 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 28 |
-
super().__init__()
|
| 29 |
-
self.device = device
|
| 30 |
-
self.features = nn.Sequential(
|
| 31 |
-
nn.Conv2d(in_shape[1], 64, kernel_size=9, stride=4, padding=2),
|
| 32 |
-
nn.ReLU(inplace=True),
|
| 33 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 34 |
-
nn.Conv2d(64, 192, kernel_size=7, padding=2),
|
| 35 |
-
nn.ReLU(inplace=True),
|
| 36 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 37 |
-
nn.Conv2d(192, 384, kernel_size=5, padding=1),
|
| 38 |
-
nn.ReLU(inplace=True),
|
| 39 |
-
nn.Conv2d(384, 256, kernel_size=3, padding=1),
|
| 40 |
-
nn.ReLU(inplace=True),
|
| 41 |
-
nn.Conv2d(256, 256, kernel_size=4, padding=1),
|
| 42 |
-
nn.ReLU(inplace=True),
|
| 43 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 44 |
-
)
|
| 45 |
-
dropout: float = prm['dropout']
|
| 46 |
-
self.avgpool = nn.AdaptiveAvgPool2d((6, 6))
|
| 47 |
-
self.classifier = nn.Sequential(
|
| 48 |
-
nn.Dropout(p=dropout),
|
| 49 |
-
nn.Linear(256 * 6 * 6, 4096),
|
| 50 |
-
nn.ReLU(inplace=True),
|
| 51 |
-
nn.Dropout(p=dropout),
|
| 52 |
-
nn.Linear(4096, 4096),
|
| 53 |
-
nn.ReLU(inplace=True),
|
| 54 |
-
nn.Linear(4096, out_shape[0]),
|
| 55 |
-
)
|
| 56 |
-
|
| 57 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 58 |
-
x = self.features(x)
|
| 59 |
-
x = self.avgpool(x)
|
| 60 |
-
x = torch.flatten(x, 1)
|
| 61 |
-
x = self.classifier(x)
|
| 62 |
-
return x
|
| 63 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/AlexNet.py
DELETED
|
@@ -1,61 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
def supported_hyperparameters():
|
| 6 |
-
return {'lr', 'momentum', 'dropout'}
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
class Net(nn.Module):
|
| 10 |
-
|
| 11 |
-
def train_setup(self, prm):
|
| 12 |
-
self.to(self.device)
|
| 13 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 14 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 15 |
-
|
| 16 |
-
def learn(self, train_data):
|
| 17 |
-
for inputs, labels in train_data:
|
| 18 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 19 |
-
self.optimizer.zero_grad()
|
| 20 |
-
outputs = self(inputs)
|
| 21 |
-
loss = self.criteria[0](outputs, labels)
|
| 22 |
-
loss.backward()
|
| 23 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 24 |
-
self.optimizer.step()
|
| 25 |
-
|
| 26 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 27 |
-
super().__init__()
|
| 28 |
-
self.device = device
|
| 29 |
-
self.features = nn.Sequential(
|
| 30 |
-
nn.Conv2d(in_shape[1], 64, kernel_size=11, stride=4, padding=2),
|
| 31 |
-
nn.ReLU(inplace=True),
|
| 32 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 33 |
-
nn.Conv2d(64, 192, kernel_size=5, padding=2),
|
| 34 |
-
nn.ReLU(inplace=True),
|
| 35 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 36 |
-
nn.Conv2d(192, 384, kernel_size=3, padding=1),
|
| 37 |
-
nn.ReLU(inplace=True),
|
| 38 |
-
nn.Conv2d(384, 256, kernel_size=3, padding=1),
|
| 39 |
-
nn.ReLU(inplace=True),
|
| 40 |
-
nn.Conv2d(256, 256, kernel_size=3, padding=1),
|
| 41 |
-
nn.ReLU(inplace=True),
|
| 42 |
-
nn.MaxPool2d(kernel_size=3, stride=2),
|
| 43 |
-
)
|
| 44 |
-
dropout: float = prm['dropout']
|
| 45 |
-
self.avgpool = nn.AdaptiveAvgPool2d((6, 6))
|
| 46 |
-
self.classifier = nn.Sequential(
|
| 47 |
-
nn.Dropout(p=dropout),
|
| 48 |
-
nn.Linear(256 * 6 * 6, 4096),
|
| 49 |
-
nn.ReLU(inplace=True),
|
| 50 |
-
nn.Dropout(p=dropout),
|
| 51 |
-
nn.Linear(4096, 4096),
|
| 52 |
-
nn.ReLU(inplace=True),
|
| 53 |
-
nn.Linear(4096, out_shape[0]),
|
| 54 |
-
)
|
| 55 |
-
|
| 56 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 57 |
-
x = self.features(x)
|
| 58 |
-
x = self.avgpool(x)
|
| 59 |
-
x = torch.flatten(x, 1)
|
| 60 |
-
x = self.classifier(x)
|
| 61 |
-
return x
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-001bf3e2-17c2-4fdf-948e-493677e58a3b.py
DELETED
|
@@ -1,134 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
def supported_hyperparameters():
|
| 6 |
-
return {'lr', 'momentum', 'dropout'}
|
| 7 |
-
|
| 8 |
-
class BagNetBottleneck(nn.Module):
|
| 9 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=4):
|
| 10 |
-
super().__init__()
|
| 11 |
-
mid_channels = out_channels // bottleneck_factor
|
| 12 |
-
|
| 13 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 14 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 15 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 16 |
-
|
| 17 |
-
@staticmethod
|
| 18 |
-
def conv1x1_block(in_channels, out_channels, activation=True):
|
| 19 |
-
return nn.Sequential(
|
| 20 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
|
| 21 |
-
nn.BatchNorm2d(out_channels) if activation else nn.Identity(),
|
| 22 |
-
nn.ReLU(inplace=True) if activation else nn.Identity(),
|
| 23 |
-
)
|
| 24 |
-
|
| 25 |
-
@staticmethod
|
| 26 |
-
def conv_block(in_channels, out_channels, kernel_size, stride):
|
| 27 |
-
padding = (kernel_size - 1) // 2
|
| 28 |
-
return nn.Sequential(
|
| 29 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 30 |
-
nn.BatchNorm2d(out_channels),
|
| 31 |
-
nn.ReLU(inplace=True),
|
| 32 |
-
)
|
| 33 |
-
|
| 34 |
-
def forward(self, x):
|
| 35 |
-
x = self.conv1(x)
|
| 36 |
-
x = self.conv2(x)
|
| 37 |
-
x = self.conv3(x)
|
| 38 |
-
return x
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
class BagNetUnit(nn.Module):
|
| 42 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 43 |
-
super().__init__()
|
| 44 |
-
self.resize_identity = (in_channels!= out_channels) or (stride!= 1)
|
| 45 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 46 |
-
|
| 47 |
-
if self.resize_identity:
|
| 48 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 49 |
-
|
| 50 |
-
self.activ = nn.ReLU(inplace=True)
|
| 51 |
-
|
| 52 |
-
@staticmethod
|
| 53 |
-
def conv1x1_block(in_channels, out_channels, activation=True):
|
| 54 |
-
return nn.Sequential(
|
| 55 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
|
| 56 |
-
nn.BatchNorm2d(out_channels) if activation else nn.Identity(),
|
| 57 |
-
nn.ReLU(inplace=True) if activation else nn.Identity(),
|
| 58 |
-
)
|
| 59 |
-
|
| 60 |
-
def forward(self, x):
|
| 61 |
-
identity = x
|
| 62 |
-
if self.resize_identity:
|
| 63 |
-
identity = self.identity_conv(x)
|
| 64 |
-
|
| 65 |
-
x = self.body(x)
|
| 66 |
-
|
| 67 |
-
if x.size(2)!= identity.size(2) or x.size(3)!= identity.size(3):
|
| 68 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 69 |
-
|
| 70 |
-
return self.activ(x + identity)
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
class Net(nn.Module):
|
| 74 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 75 |
-
super().__init__()
|
| 76 |
-
self.device = device
|
| 77 |
-
channel_number = in_shape[1]
|
| 78 |
-
image_size = in_shape[2]
|
| 79 |
-
class_number = out_shape[0]
|
| 80 |
-
learning_rate = prm['lr']
|
| 81 |
-
momentum = prm['momentum']
|
| 82 |
-
dropout = prm['dropout']
|
| 83 |
-
|
| 84 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 85 |
-
self.in_size = image_size
|
| 86 |
-
self.num_classes = class_number
|
| 87 |
-
|
| 88 |
-
self.features = nn.Sequential(
|
| 89 |
-
nn.Conv2d(channel_number, 64, kernel_size=7, stride=2, padding=3, bias=False),
|
| 90 |
-
nn.BatchNorm2d(64),
|
| 91 |
-
nn.ReLU(inplace=True),
|
| 92 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
|
| 93 |
-
)
|
| 94 |
-
|
| 95 |
-
in_channels = 64
|
| 96 |
-
for i, stage_channels in enumerate(self.channels):
|
| 97 |
-
stage = nn.Sequential()
|
| 98 |
-
for j, out_channels in enumerate(stage_channels):
|
| 99 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 100 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 101 |
-
in_channels = out_channels
|
| 102 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 103 |
-
|
| 104 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 105 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 106 |
-
|
| 107 |
-
self.learning_rate = learning_rate
|
| 108 |
-
self.momentum = momentum
|
| 109 |
-
self.dropout = dropout
|
| 110 |
-
|
| 111 |
-
def forward(self, x):
|
| 112 |
-
x = self.features(x)
|
| 113 |
-
x = torch.flatten(x, 1)
|
| 114 |
-
return self.output(x)
|
| 115 |
-
|
| 116 |
-
def train_setup(self, prm):
|
| 117 |
-
self.to(self.device)
|
| 118 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 119 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 120 |
-
|
| 121 |
-
if self.dropout > 0:
|
| 122 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 123 |
-
|
| 124 |
-
def learn(self, train_data):
|
| 125 |
-
self.train()
|
| 126 |
-
for inputs, labels in train_data:
|
| 127 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 128 |
-
self.optimizer.zero_grad()
|
| 129 |
-
outputs = self(inputs)
|
| 130 |
-
loss = self.criteria(outputs, labels)
|
| 131 |
-
loss.backward()
|
| 132 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 133 |
-
self.optimizer.step()
|
| 134 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-560341b5-15a8-4829-a8ac-ea4c7391a950.py
DELETED
|
@@ -1,129 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class BagNetBottleneck(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=8): # Changed 4 to 8
|
| 12 |
-
super().__init__()
|
| 13 |
-
mid_channels = out_channels // bottleneck_factor
|
| 14 |
-
|
| 15 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 16 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 17 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 18 |
-
|
| 19 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 20 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 21 |
-
if activation:
|
| 22 |
-
layers.append(nn.ReLU(inplace=True))
|
| 23 |
-
return nn.Sequential(*layers)
|
| 24 |
-
|
| 25 |
-
def conv_block(self, in_channels, out_channels, kernel_size, stride):
|
| 26 |
-
padding = (kernel_size - 1) // 2
|
| 27 |
-
return nn.Sequential(
|
| 28 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 29 |
-
nn.BatchNorm2d(out_channels),
|
| 30 |
-
nn.ReLU(inplace=True),
|
| 31 |
-
)
|
| 32 |
-
|
| 33 |
-
def forward(self, x):
|
| 34 |
-
x = self.conv1(x)
|
| 35 |
-
x = self.conv2(x)
|
| 36 |
-
x = self.conv3(x)
|
| 37 |
-
return x
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
class BagNetUnit(nn.Module):
|
| 41 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 42 |
-
super().__init__()
|
| 43 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 44 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 45 |
-
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 51 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 52 |
-
if activation:
|
| 53 |
-
layers.append(nn.ReLU(inplace=True))
|
| 54 |
-
return nn.Sequential(*layers)
|
| 55 |
-
|
| 56 |
-
def forward(self, x):
|
| 57 |
-
identity = x
|
| 58 |
-
if self.resize_identity:
|
| 59 |
-
identity = self.identity_conv(x)
|
| 60 |
-
|
| 61 |
-
x = self.body(x)
|
| 62 |
-
|
| 63 |
-
if x.size(2) != identity.size(2) or x.size(3) != identity.size(3):
|
| 64 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 65 |
-
|
| 66 |
-
return self.activ(x + identity)
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
class Net(nn.Module):
|
| 70 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 71 |
-
super().__init__()
|
| 72 |
-
self.device = device
|
| 73 |
-
channel_number = in_shape[1]
|
| 74 |
-
image_size = in_shape[2]
|
| 75 |
-
class_number = out_shape[0]
|
| 76 |
-
learning_rate = prm['lr']
|
| 77 |
-
momentum = prm['momentum']
|
| 78 |
-
dropout = prm['dropout']
|
| 79 |
-
|
| 80 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 81 |
-
self.in_size = image_size
|
| 82 |
-
self.num_classes = class_number
|
| 83 |
-
|
| 84 |
-
self.features = nn.Sequential(
|
| 85 |
-
nn.Conv2d(channel_number, 64, kernel_size=7, stride=2, padding=3, bias=False),
|
| 86 |
-
nn.BatchNorm2d(64),
|
| 87 |
-
nn.ReLU(inplace=True),
|
| 88 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
|
| 89 |
-
)
|
| 90 |
-
|
| 91 |
-
in_channels = 64
|
| 92 |
-
for i, stage_channels in enumerate(self.channels):
|
| 93 |
-
stage = nn.Sequential()
|
| 94 |
-
for j, out_channels in enumerate(stage_channels):
|
| 95 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 96 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 97 |
-
in_channels = out_channels
|
| 98 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 99 |
-
|
| 100 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 101 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 102 |
-
|
| 103 |
-
self.learning_rate = learning_rate
|
| 104 |
-
self.momentum = momentum
|
| 105 |
-
self.dropout = dropout
|
| 106 |
-
|
| 107 |
-
def forward(self, x):
|
| 108 |
-
x = self.features(x)
|
| 109 |
-
x = torch.flatten(x, 1)
|
| 110 |
-
return self.output(x)
|
| 111 |
-
|
| 112 |
-
def train_setup(self, prm):
|
| 113 |
-
self.to(self.device)
|
| 114 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 115 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 116 |
-
|
| 117 |
-
if self.dropout > 0:
|
| 118 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 119 |
-
|
| 120 |
-
def learn(self, train_data):
|
| 121 |
-
self.train()
|
| 122 |
-
for inputs, labels in train_data:
|
| 123 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 124 |
-
self.optimizer.zero_grad()
|
| 125 |
-
outputs = self(inputs)
|
| 126 |
-
loss = self.criteria(outputs, labels)
|
| 127 |
-
loss.backward()
|
| 128 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 129 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-6ecd3fc7-5ce2-4876-86a6-5e8250c43d78.py
DELETED
|
@@ -1,129 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class BagNetBottleneck(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=4):
|
| 12 |
-
super().__init__()
|
| 13 |
-
mid_channels = out_channels // bottleneck_factor
|
| 14 |
-
|
| 15 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 16 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 17 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 18 |
-
|
| 19 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 20 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 21 |
-
if activation:
|
| 22 |
-
layers.append(nn.ReLU(inplace=True))
|
| 23 |
-
return nn.Sequential(*layers)
|
| 24 |
-
|
| 25 |
-
def conv_block(self, in_channels, out_channels, kernel_size, stride):
|
| 26 |
-
padding = (kernel_size - 1) // 2
|
| 27 |
-
return nn.Sequential(
|
| 28 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 29 |
-
nn.BatchNorm2d(out_channels),
|
| 30 |
-
nn.ReLU(inplace=True),
|
| 31 |
-
)
|
| 32 |
-
|
| 33 |
-
def forward(self, x):
|
| 34 |
-
x = self.conv1(x)
|
| 35 |
-
x = self.conv2(x)
|
| 36 |
-
x = self.conv3(x)
|
| 37 |
-
return x
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
class BagNetUnit(nn.Module):
|
| 41 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 42 |
-
super().__init__()
|
| 43 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 44 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 45 |
-
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 51 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 52 |
-
if activation:
|
| 53 |
-
layers.append(nn.ReLU(inplace=True))
|
| 54 |
-
return nn.Sequential(*layers)
|
| 55 |
-
|
| 56 |
-
def forward(self, x):
|
| 57 |
-
identity = x
|
| 58 |
-
if self.resize_identity:
|
| 59 |
-
identity = self.identity_conv(x)
|
| 60 |
-
|
| 61 |
-
x = self.body(x)
|
| 62 |
-
|
| 63 |
-
if x.size(2) != identity.size(2) or x.size(3) != identity.size(3):
|
| 64 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 65 |
-
|
| 66 |
-
return self.activ(x + identity)
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
class Net(nn.Module):
|
| 70 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 71 |
-
super().__init__()
|
| 72 |
-
self.device = device
|
| 73 |
-
channel_number = in_shape[1]
|
| 74 |
-
image_size = in_shape[2]
|
| 75 |
-
class_number = out_shape[0]
|
| 76 |
-
learning_rate = prm['lr']
|
| 77 |
-
momentum = prm['momentum']
|
| 78 |
-
dropout = prm['dropout']
|
| 79 |
-
|
| 80 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 81 |
-
self.in_size = image_size
|
| 82 |
-
self.num_classes = class_number
|
| 83 |
-
|
| 84 |
-
self.features = nn.Sequential(
|
| 85 |
-
nn.Conv2d(channel_number, 64, kernel_size=5, stride=2, padding=3, bias=False), # Changed kernel_size from 7 to 5
|
| 86 |
-
nn.BatchNorm2d(64),
|
| 87 |
-
nn.ReLU(inplace=True),
|
| 88 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
|
| 89 |
-
)
|
| 90 |
-
|
| 91 |
-
in_channels = 64
|
| 92 |
-
for i, stage_channels in enumerate(self.channels):
|
| 93 |
-
stage = nn.Sequential()
|
| 94 |
-
for j, out_channels in enumerate(stage_channels):
|
| 95 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 96 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 97 |
-
in_channels = out_channels
|
| 98 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 99 |
-
|
| 100 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 101 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 102 |
-
|
| 103 |
-
self.learning_rate = learning_rate
|
| 104 |
-
self.momentum = momentum
|
| 105 |
-
self.dropout = dropout
|
| 106 |
-
|
| 107 |
-
def forward(self, x):
|
| 108 |
-
x = self.features(x)
|
| 109 |
-
x = torch.flatten(x, 1)
|
| 110 |
-
return self.output(x)
|
| 111 |
-
|
| 112 |
-
def train_setup(self, prm):
|
| 113 |
-
self.to(self.device)
|
| 114 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 115 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 116 |
-
|
| 117 |
-
if self.dropout > 0:
|
| 118 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 119 |
-
|
| 120 |
-
def learn(self, train_data):
|
| 121 |
-
self.train()
|
| 122 |
-
for inputs, labels in train_data:
|
| 123 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 124 |
-
self.optimizer.zero_grad()
|
| 125 |
-
outputs = self(inputs)
|
| 126 |
-
loss = self.criteria(outputs, labels)
|
| 127 |
-
loss.backward()
|
| 128 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 129 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-7e541be1-6b60-445d-bbbf-3b655eeefc9a.py
DELETED
|
@@ -1,129 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class BagNetBottleneck(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=6):
|
| 12 |
-
super().__init__()
|
| 13 |
-
mid_channels = out_channels // bottleneck_factor
|
| 14 |
-
|
| 15 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 16 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 17 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 18 |
-
|
| 19 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 20 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 21 |
-
if activation:
|
| 22 |
-
layers.append(nn.ReLU(inplace=True))
|
| 23 |
-
return nn.Sequential(*layers)
|
| 24 |
-
|
| 25 |
-
def conv_block(self, in_channels, out_channels, kernel_size, stride):
|
| 26 |
-
padding = (kernel_size - 1) // 2
|
| 27 |
-
return nn.Sequential(
|
| 28 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 29 |
-
nn.BatchNorm2d(out_channels),
|
| 30 |
-
nn.ReLU(inplace=True),
|
| 31 |
-
)
|
| 32 |
-
|
| 33 |
-
def forward(self, x):
|
| 34 |
-
x = self.conv1(x)
|
| 35 |
-
x = self.conv2(x)
|
| 36 |
-
x = self.conv3(x)
|
| 37 |
-
return x
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
class BagNetUnit(nn.Module):
|
| 41 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 42 |
-
super().__init__()
|
| 43 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 44 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 45 |
-
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 51 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 52 |
-
if activation:
|
| 53 |
-
layers.append(nn.ReLU(inplace=True))
|
| 54 |
-
return nn.Sequential(*layers)
|
| 55 |
-
|
| 56 |
-
def forward(self, x):
|
| 57 |
-
identity = x
|
| 58 |
-
if self.resize_identity:
|
| 59 |
-
identity = self.identity_conv(x)
|
| 60 |
-
|
| 61 |
-
x = self.body(x)
|
| 62 |
-
|
| 63 |
-
if x.size(2) != identity.size(2) or x.size(3) != identity.size(3):
|
| 64 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 65 |
-
|
| 66 |
-
return self.activ(x + identity)
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
class Net(nn.Module):
|
| 70 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 71 |
-
super().__init__()
|
| 72 |
-
self.device = device
|
| 73 |
-
channel_number = in_shape[1]
|
| 74 |
-
image_size = in_shape[2]
|
| 75 |
-
class_number = out_shape[0]
|
| 76 |
-
learning_rate = prm['lr']
|
| 77 |
-
momentum = prm['momentum']
|
| 78 |
-
dropout = prm['dropout']
|
| 79 |
-
|
| 80 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 81 |
-
self.in_size = image_size
|
| 82 |
-
self.num_classes = class_number
|
| 83 |
-
|
| 84 |
-
self.features = nn.Sequential(
|
| 85 |
-
nn.Conv2d(channel_number, 64, kernel_size=5, stride=2, padding=2, bias=False),
|
| 86 |
-
nn.BatchNorm2d(64),
|
| 87 |
-
nn.ReLU(inplace=True),
|
| 88 |
-
nn.MaxPool2d(kernel_size=2, stride=2, padding=1),
|
| 89 |
-
)
|
| 90 |
-
|
| 91 |
-
in_channels = 64
|
| 92 |
-
for i, stage_channels in enumerate(self.channels):
|
| 93 |
-
stage = nn.Sequential()
|
| 94 |
-
for j, out_channels in enumerate(stage_channels):
|
| 95 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 96 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 97 |
-
in_channels = out_channels
|
| 98 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 99 |
-
|
| 100 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 101 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 102 |
-
|
| 103 |
-
self.learning_rate = learning_rate
|
| 104 |
-
self.momentum = momentum
|
| 105 |
-
self.dropout = dropout
|
| 106 |
-
|
| 107 |
-
def forward(self, x):
|
| 108 |
-
x = self.features(x)
|
| 109 |
-
x = torch.flatten(x, 1)
|
| 110 |
-
return self.output(x)
|
| 111 |
-
|
| 112 |
-
def train_setup(self, prm):
|
| 113 |
-
self.to(self.device)
|
| 114 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 115 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 116 |
-
|
| 117 |
-
if self.dropout > 0:
|
| 118 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 119 |
-
|
| 120 |
-
def learn(self, train_data):
|
| 121 |
-
self.train()
|
| 122 |
-
for inputs, labels in train_data:
|
| 123 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 124 |
-
self.optimizer.zero_grad()
|
| 125 |
-
outputs = self(inputs)
|
| 126 |
-
loss = self.criteria(outputs, labels)
|
| 127 |
-
loss.backward()
|
| 128 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 129 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-7ebe6562-46c6-4406-96a2-bf3914ac8516.py
DELETED
|
@@ -1,139 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
class SupportedHyperparameters():
|
| 7 |
-
def __init__(self):
|
| 8 |
-
self.hyperparameters = {'lr','momentum', 'dropout'}
|
| 9 |
-
|
| 10 |
-
def check_hyperparameters(self, param):
|
| 11 |
-
return param in self.hyperparameters
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
class BagNetBottleneck(nn.Module):
|
| 15 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=4):
|
| 16 |
-
super().__init__()
|
| 17 |
-
mid_channels = out_channels // bottleneck_factor
|
| 18 |
-
|
| 19 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 20 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 21 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 22 |
-
|
| 23 |
-
@staticmethod
|
| 24 |
-
def conv1x1_block(in_channels, out_channels, activation=True):
|
| 25 |
-
return nn.Sequential(
|
| 26 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
|
| 27 |
-
nn.BatchNorm2d(out_channels) if activation else nn.Identity(),
|
| 28 |
-
nn.ReLU(inplace=True) if activation else nn.Identity(),
|
| 29 |
-
)
|
| 30 |
-
|
| 31 |
-
@staticmethod
|
| 32 |
-
def conv_block(in_channels, out_channels, kernel_size, stride):
|
| 33 |
-
padding = (kernel_size - 1) // 2
|
| 34 |
-
return nn.Sequential(
|
| 35 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 36 |
-
nn.BatchNorm2d(out_channels),
|
| 37 |
-
nn.ReLU(inplace=True),
|
| 38 |
-
)
|
| 39 |
-
|
| 40 |
-
def forward(self, x):
|
| 41 |
-
x = self.conv1(x)
|
| 42 |
-
x = self.conv2(x)
|
| 43 |
-
x = self.conv3(x)
|
| 44 |
-
return x
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
class BagNetUnit(nn.Module):
|
| 48 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 49 |
-
super().__init__()
|
| 50 |
-
self.resize_identity = (in_channels!= out_channels) or (stride!= 1)
|
| 51 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 52 |
-
|
| 53 |
-
if self.resize_identity:
|
| 54 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 55 |
-
|
| 56 |
-
self.activ = nn.ReLU(inplace=True)
|
| 57 |
-
|
| 58 |
-
@staticmethod
|
| 59 |
-
def conv1x1_block(in_channels, out_channels, activation=True):
|
| 60 |
-
return nn.Sequential(
|
| 61 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
|
| 62 |
-
nn.BatchNorm2d(out_channels) if activation else nn.Identity(),
|
| 63 |
-
nn.ReLU(inplace=True) if activation else nn.Identity(),
|
| 64 |
-
)
|
| 65 |
-
|
| 66 |
-
def forward(self, x):
|
| 67 |
-
identity = x
|
| 68 |
-
if self.resize_identity:
|
| 69 |
-
identity = self.identity_conv(x)
|
| 70 |
-
|
| 71 |
-
x = self.body(x)
|
| 72 |
-
|
| 73 |
-
if x.size(2)!= identity.size(2) or x.size(3)!= identity.size(3):
|
| 74 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 75 |
-
|
| 76 |
-
return self.activ(x + identity)
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
class Net(nn.Module):
|
| 80 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 81 |
-
super().__init__()
|
| 82 |
-
self.device = device
|
| 83 |
-
channel_number = in_shape[1]
|
| 84 |
-
image_size = in_shape[2]
|
| 85 |
-
class_number = out_shape[0]
|
| 86 |
-
learning_rate = prm['lr']
|
| 87 |
-
momentum = prm['momentum']
|
| 88 |
-
dropout = prm['dropout']
|
| 89 |
-
|
| 90 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 91 |
-
self.in_size = image_size
|
| 92 |
-
self.num_classes = class_number
|
| 93 |
-
|
| 94 |
-
self.features = nn.Sequential(
|
| 95 |
-
nn.Conv2d(channel_number, 64, kernel_size=7, stride=2, padding=3, bias=False),
|
| 96 |
-
nn.BatchNorm2d(64),
|
| 97 |
-
nn.ReLU(inplace=True),
|
| 98 |
-
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
|
| 99 |
-
)
|
| 100 |
-
|
| 101 |
-
in_channels = 64
|
| 102 |
-
for i, stage_channels in enumerate(self.channels):
|
| 103 |
-
stage = nn.Sequential()
|
| 104 |
-
for j, out_channels in enumerate(stage_channels):
|
| 105 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 106 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 107 |
-
in_channels = out_channels
|
| 108 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 109 |
-
|
| 110 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 111 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 112 |
-
|
| 113 |
-
self.learning_rate = learning_rate
|
| 114 |
-
self.momentum = momentum
|
| 115 |
-
self.dropout = dropout
|
| 116 |
-
|
| 117 |
-
def forward(self, x):
|
| 118 |
-
x = self.features(x)
|
| 119 |
-
x = torch.flatten(x, 1)
|
| 120 |
-
return self.output(x)
|
| 121 |
-
|
| 122 |
-
def train_setup(self, prm):
|
| 123 |
-
self.to(self.device)
|
| 124 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 125 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 126 |
-
|
| 127 |
-
if self.dropout > 0:
|
| 128 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 129 |
-
|
| 130 |
-
def learn(self, train_data):
|
| 131 |
-
self.train()
|
| 132 |
-
for inputs, labels in train_data:
|
| 133 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 134 |
-
self.optimizer.zero_grad()
|
| 135 |
-
outputs = self(inputs)
|
| 136 |
-
loss = self.criteria(outputs, labels)
|
| 137 |
-
loss.backward()
|
| 138 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 139 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BagNet-7f792262-31cf-477e-a78a-3494c122332d.py
DELETED
|
@@ -1,129 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class BagNetBottleneck(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride, bottleneck_factor=4):
|
| 12 |
-
super().__init__()
|
| 13 |
-
mid_channels = out_channels // bottleneck_factor
|
| 14 |
-
|
| 15 |
-
self.conv1 = self.conv1x1_block(in_channels, mid_channels)
|
| 16 |
-
self.conv2 = self.conv_block(mid_channels, mid_channels, kernel_size, stride)
|
| 17 |
-
self.conv3 = self.conv1x1_block(mid_channels, out_channels, activation=False)
|
| 18 |
-
|
| 19 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 20 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 21 |
-
if activation:
|
| 22 |
-
layers.append(nn.ReLU(inplace=True))
|
| 23 |
-
return nn.Sequential(*layers)
|
| 24 |
-
|
| 25 |
-
def conv_block(self, in_channels, out_channels, kernel_size, stride):
|
| 26 |
-
padding = (kernel_size - 1) // 2
|
| 27 |
-
return nn.Sequential(
|
| 28 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding, bias=False),
|
| 29 |
-
nn.BatchNorm2d(out_channels),
|
| 30 |
-
nn.ReLU(inplace=True),
|
| 31 |
-
)
|
| 32 |
-
|
| 33 |
-
def forward(self, x):
|
| 34 |
-
x = self.conv1(x)
|
| 35 |
-
x = self.conv2(x)
|
| 36 |
-
x = self.conv3(x)
|
| 37 |
-
return x
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
class BagNetUnit(nn.Module):
|
| 41 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride):
|
| 42 |
-
super().__init__()
|
| 43 |
-
self.resize_identity = (in_channels != out_channels) or (stride != 1)
|
| 44 |
-
self.body = BagNetBottleneck(in_channels, out_channels, kernel_size, stride)
|
| 45 |
-
|
| 46 |
-
if self.resize_identity:
|
| 47 |
-
self.identity_conv = self.conv1x1_block(in_channels, out_channels, activation=False)
|
| 48 |
-
self.activ = nn.ReLU(inplace=True)
|
| 49 |
-
|
| 50 |
-
def conv1x1_block(self, in_channels, out_channels, activation=True):
|
| 51 |
-
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)]
|
| 52 |
-
if activation:
|
| 53 |
-
layers.append(nn.ReLU(inplace=True))
|
| 54 |
-
return nn.Sequential(*layers)
|
| 55 |
-
|
| 56 |
-
def forward(self, x):
|
| 57 |
-
identity = x
|
| 58 |
-
if self.resize_identity:
|
| 59 |
-
identity = self.identity_conv(x)
|
| 60 |
-
|
| 61 |
-
x = self.body(x)
|
| 62 |
-
|
| 63 |
-
if x.size(2) != identity.size(2) or x.size(3) != identity.size(3):
|
| 64 |
-
identity = nn.functional.interpolate(identity, size=(x.size(2), x.size(3)), mode='bilinear', align_corners=False)
|
| 65 |
-
|
| 66 |
-
return self.activ(x + identity)
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
class Net(nn.Module):
|
| 70 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 71 |
-
super().__init__()
|
| 72 |
-
self.device = device
|
| 73 |
-
channel_number = in_shape[1]
|
| 74 |
-
image_size = in_shape[2]
|
| 75 |
-
class_number = out_shape[0]
|
| 76 |
-
learning_rate = prm['lr']
|
| 77 |
-
momentum = prm['momentum']
|
| 78 |
-
dropout = prm['dropout']
|
| 79 |
-
|
| 80 |
-
self.channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 81 |
-
self.in_size = image_size
|
| 82 |
-
self.num_classes = class_number
|
| 83 |
-
|
| 84 |
-
self.features = nn.Sequential(
|
| 85 |
-
nn.Conv2d(channel_number, 64, kernel_size=5, stride=3, padding=2), # Changed kernel_size from 7 to 5
|
| 86 |
-
nn.BatchNorm2d(64),
|
| 87 |
-
nn.ReLU(inplace=True),
|
| 88 |
-
nn.MaxPool2d(kernel_size=5, stride=4, padding=0), # Changed padding from 1 to 0
|
| 89 |
-
)
|
| 90 |
-
|
| 91 |
-
in_channels = 64
|
| 92 |
-
for i, stage_channels in enumerate(self.channels):
|
| 93 |
-
stage = nn.Sequential()
|
| 94 |
-
for j, out_channels in enumerate(stage_channels):
|
| 95 |
-
stride = 2 if (j == 0 and i > 0) else 1
|
| 96 |
-
stage.add_module(f"unit{j + 1}", BagNetUnit(in_channels, out_channels, kernel_size=3, stride=stride))
|
| 97 |
-
in_channels = out_channels
|
| 98 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 99 |
-
|
| 100 |
-
self.features.add_module("final_pool", nn.AdaptiveAvgPool2d(1))
|
| 101 |
-
self.output = nn.Linear(in_channels, self.num_classes)
|
| 102 |
-
|
| 103 |
-
self.learning_rate = learning_rate
|
| 104 |
-
self.momentum = momentum
|
| 105 |
-
self.dropout = dropout
|
| 106 |
-
|
| 107 |
-
def forward(self, x):
|
| 108 |
-
x = self.features(x)
|
| 109 |
-
x = torch.flatten(x, 1)
|
| 110 |
-
return self.output(x)
|
| 111 |
-
|
| 112 |
-
def train_setup(self, prm):
|
| 113 |
-
self.to(self.device)
|
| 114 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 115 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'],)
|
| 116 |
-
|
| 117 |
-
if self.dropout > 0:
|
| 118 |
-
self.dropout_layer = nn.Dropout(self.dropout)
|
| 119 |
-
|
| 120 |
-
def learn(self, train_data):
|
| 121 |
-
self.train()
|
| 122 |
-
for inputs, labels in train_data:
|
| 123 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 124 |
-
self.optimizer.zero_grad()
|
| 125 |
-
outputs = self(inputs)
|
| 126 |
-
loss = self.criteria(outputs, labels)
|
| 127 |
-
loss.backward()
|
| 128 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 129 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BayesianNet-024b0436-9ad5-4a1f-86d0-946e577ffc2d.py
DELETED
|
@@ -1,244 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
from torch.nn import Parameter
|
| 5 |
-
|
| 6 |
-
def calculate_kl(mu_q, sig_q, mu_p, sig_p, eps=1e-8):
|
| 7 |
-
kl = 0.5 * (2 * torch.log(sig_p / sig_q) - 1 + (sig_q / sig_p).pow(2) + ((mu_p - mu_q) / sig_p).pow(2)).sum() + eps
|
| 8 |
-
return kl
|
| 9 |
-
|
| 10 |
-
class ModuleWrapper(nn.Module):
|
| 11 |
-
def __init__(self):
|
| 12 |
-
super(ModuleWrapper, self).__init__()
|
| 13 |
-
|
| 14 |
-
def set_flag(self, flag_name, value):
|
| 15 |
-
setattr(self, flag_name, value)
|
| 16 |
-
for m in self.children():
|
| 17 |
-
if hasattr(m,'set_flag'):
|
| 18 |
-
m.set_flag(flag_name, value)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
for module in self.children():
|
| 22 |
-
x = module(x)
|
| 23 |
-
|
| 24 |
-
kl = 0.0
|
| 25 |
-
for module in self.modules():
|
| 26 |
-
if hasattr(module, 'kl_loss'):
|
| 27 |
-
kl = kl + module.kl_loss()
|
| 28 |
-
|
| 29 |
-
return x, kl
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
class FlattenLayer(ModuleWrapper):
|
| 33 |
-
|
| 34 |
-
def __init__(self, num_features):
|
| 35 |
-
super(FlattenLayer, self).__init__()
|
| 36 |
-
self.num_features = num_features
|
| 37 |
-
|
| 38 |
-
def forward(self, x):
|
| 39 |
-
return x.view(-1, self.num_features)
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
class BBBLinear(ModuleWrapper):
|
| 43 |
-
|
| 44 |
-
def __init__(self, in_features, out_features, bias=True, priors=None):
|
| 45 |
-
super(BBBLinear, self).__init__()
|
| 46 |
-
self.in_features = in_features
|
| 47 |
-
self.out_features = out_features
|
| 48 |
-
self.use_bias = bias
|
| 49 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 50 |
-
|
| 51 |
-
if priors is None:
|
| 52 |
-
priors = {
|
| 53 |
-
'prior_mu': 0, # Randomly changed from 0.5
|
| 54 |
-
'prior_sigma': 0.1,
|
| 55 |
-
'posterior_mu_initial': (0, 0.1),
|
| 56 |
-
'posterior_rho_initial': (-5, 0.1),
|
| 57 |
-
}
|
| 58 |
-
self.prior_mu = priors['prior_mu']
|
| 59 |
-
self.prior_sigma = priors['prior_sigma']
|
| 60 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 61 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 62 |
-
|
| 63 |
-
self.W_mu = Parameter(torch.Tensor(out_features, in_features))
|
| 64 |
-
self.W_rho = Parameter(torch.Tensor(out_features, in_features))
|
| 65 |
-
if self.use_bias:
|
| 66 |
-
self.bias_mu = Parameter(torch.Tensor(out_features))
|
| 67 |
-
self.bias_rho = Parameter(torch.Tensor(out_features))
|
| 68 |
-
else:
|
| 69 |
-
self.register_parameter('bias_mu', None)
|
| 70 |
-
self.register_parameter('bias_rho', None)
|
| 71 |
-
|
| 72 |
-
self.reset_parameters()
|
| 73 |
-
|
| 74 |
-
def reset_parameters(self):
|
| 75 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 76 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 77 |
-
|
| 78 |
-
if self.use_bias:
|
| 79 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 80 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 81 |
-
|
| 82 |
-
def forward(self, x, sample=True):
|
| 83 |
-
|
| 84 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 85 |
-
if self.use_bias:
|
| 86 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 87 |
-
bias_var = self.bias_sigma ** 2
|
| 88 |
-
else:
|
| 89 |
-
self.bias_sigma = bias_var = None
|
| 90 |
-
|
| 91 |
-
act_mu = F.linear(x, self.W_mu, self.bias_mu)
|
| 92 |
-
act_var = 1e-16 + F.linear(x ** 2, self.W_sigma ** 2, bias_var)
|
| 93 |
-
act_std = torch.sqrt(act_var)
|
| 94 |
-
|
| 95 |
-
if self.training or sample:
|
| 96 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 97 |
-
return act_mu + act_std * eps
|
| 98 |
-
else:
|
| 99 |
-
return act_mu
|
| 100 |
-
|
| 101 |
-
def kl_loss(self):
|
| 102 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 103 |
-
if self.use_bias:
|
| 104 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 105 |
-
return kl
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
class BBBConv2d(ModuleWrapper):
|
| 109 |
-
|
| 110 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
|
| 111 |
-
padding=0, dilation=1, bias=True, priors=None):
|
| 112 |
-
super(BBBConv2d, self).__init__()
|
| 113 |
-
self.in_channels = in_channels
|
| 114 |
-
self.out_channels = out_channels
|
| 115 |
-
self.kernel_size = (kernel_size, kernel_size) # Randomly changed from 5 to 7
|
| 116 |
-
self.stride = stride
|
| 117 |
-
self.padding = padding
|
| 118 |
-
self.dilation = dilation
|
| 119 |
-
self.groups = 1
|
| 120 |
-
self.use_bias = bias
|
| 121 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 122 |
-
|
| 123 |
-
if priors is None:
|
| 124 |
-
priors = {
|
| 125 |
-
'prior_mu': 0, # Randomly changed from 0.5
|
| 126 |
-
'prior_sigma': 0.1,
|
| 127 |
-
'posterior_mu_initial': (0, 0.1),
|
| 128 |
-
'posterior_rho_initial': (-5, 0.1),
|
| 129 |
-
}
|
| 130 |
-
self.prior_mu = priors['prior_mu']
|
| 131 |
-
self.prior_sigma = priors['prior_sigma']
|
| 132 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 133 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 134 |
-
|
| 135 |
-
self.W_mu = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 136 |
-
self.W_rho = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 137 |
-
if self.use_bias:
|
| 138 |
-
self.bias_mu = Parameter(torch.Tensor(out_channels))
|
| 139 |
-
self.bias_rho = Parameter(torch.Tensor(out_channels))
|
| 140 |
-
else:
|
| 141 |
-
self.register_parameter('bias_mu', None)
|
| 142 |
-
self.register_parameter('bias_rho', None)
|
| 143 |
-
|
| 144 |
-
self.reset_parameters()
|
| 145 |
-
|
| 146 |
-
def reset_parameters(self):
|
| 147 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 148 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 149 |
-
|
| 150 |
-
if self.use_bias:
|
| 151 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 152 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 153 |
-
|
| 154 |
-
def forward(self, x, sample=True):
|
| 155 |
-
|
| 156 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 157 |
-
if self.use_bias:
|
| 158 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 159 |
-
bias_var = self.bias_sigma ** 2
|
| 160 |
-
else:
|
| 161 |
-
self.bias_sigma = bias_var = None
|
| 162 |
-
|
| 163 |
-
act_mu = F.conv2d(
|
| 164 |
-
x, self.W_mu, self.bias_mu, self.stride, self.padding, self.dilation, self.groups)
|
| 165 |
-
act_var = 1e-16 + F.conv2d(
|
| 166 |
-
x ** 2, self.W_sigma ** 2, bias_var, self.stride, self.padding, self.dilation, self.groups)
|
| 167 |
-
act_std = torch.sqrt(act_var)
|
| 168 |
-
|
| 169 |
-
if self.training or sample:
|
| 170 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 171 |
-
return act_mu + act_std * eps
|
| 172 |
-
else:
|
| 173 |
-
return act_mu
|
| 174 |
-
|
| 175 |
-
def kl_loss(self):
|
| 176 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 177 |
-
if self.use_bias:
|
| 178 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 179 |
-
return kl
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
def supported_hyperparameters():
|
| 183 |
-
return {'lr','momentum', 'kernel_size'} # Randomly changed from'momentum' to 'kernel_size'
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
class Net(ModuleWrapper):
|
| 187 |
-
|
| 188 |
-
def train_setup(self, prm):
|
| 189 |
-
self.to(self.device)
|
| 190 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 191 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 192 |
-
if 'kernel_size' in prm:
|
| 193 |
-
self.kernal_size = prm['kernel_size'] # Randomly changed from default
|
| 194 |
-
|
| 195 |
-
def learn(self, train_data):
|
| 196 |
-
for inputs, labels in train_data:
|
| 197 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 198 |
-
self.optimizer.zero_grad()
|
| 199 |
-
outputs = self(inputs)
|
| 200 |
-
loss = self.criteria[0](outputs, labels)
|
| 201 |
-
loss.backward()
|
| 202 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 203 |
-
self.optimizer.step()
|
| 204 |
-
|
| 205 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 206 |
-
super(Net, self).__init__()
|
| 207 |
-
self.device = device
|
| 208 |
-
inputs: int = in_shape[1]
|
| 209 |
-
outputs: int = out_shape[0]
|
| 210 |
-
self.num_classes = outputs
|
| 211 |
-
self.priors = {
|
| 212 |
-
'prior_mu': 0,
|
| 213 |
-
'prior_sigma': 0.1,
|
| 214 |
-
'posterior_mu_initial': (0, 0.1),
|
| 215 |
-
'posterior_rho_initial': (-5, 0.1),
|
| 216 |
-
}
|
| 217 |
-
self.act = nn.Softplus
|
| 218 |
-
|
| 219 |
-
self.conv1 = BBBConv2d(inputs, 64, 11, stride=4, padding=5, bias=True, priors=self.priors)
|
| 220 |
-
self.act1 = self.act()
|
| 221 |
-
self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 222 |
-
|
| 223 |
-
self.conv2 = BBBConv2d(64, 192, 5, padding=2, bias=True, priors=self.priors) # Randomly changed from 5 to 7
|
| 224 |
-
self.act2 = self.act()
|
| 225 |
-
self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 226 |
-
|
| 227 |
-
self.conv3 = BBBConv2d(192, 384, 3, padding=1, bias=True, priors=self.priors)
|
| 228 |
-
self.act3 = self.act()
|
| 229 |
-
|
| 230 |
-
self.conv4 = BBBConv2d(384, 256, 3, padding=1, bias=True, priors=self.priors)
|
| 231 |
-
self.act4 = self.act()
|
| 232 |
-
|
| 233 |
-
self.conv5 = BBBConv2d(256, 128, 3, padding=1, bias=True, priors=self.priors)
|
| 234 |
-
self.act5 = self.act()
|
| 235 |
-
self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 236 |
-
|
| 237 |
-
self.flatten = FlattenLayer(1 * 1 * 128)
|
| 238 |
-
self.classifier = BBBLinear(1 * 1 * 128, outputs, bias=True, priors=self.priors)
|
| 239 |
-
|
| 240 |
-
def __call__(self, *args):
|
| 241 |
-
return super().__call__(*args)[0]
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
### END OF REPLY
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BayesianNet-0901ac22-d7f5-4deb-94c9-970e7955bd68.py
DELETED
|
@@ -1,242 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
from torch.nn import Parameter
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
def calculate_kl(mu_q, sig_q, mu_p, sig_p):
|
| 9 |
-
kl = 0.5 * (2 * torch.log(sig_p / sig_q) - 1 + (sig_q / sig_p).pow(2) + ((mu_p - mu_q) / sig_p).pow(2)).sum()
|
| 10 |
-
return kl
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
class ModuleWrapper(nn.Module):
|
| 14 |
-
def __init__(self):
|
| 15 |
-
super(ModuleWrapper, self).__init__()
|
| 16 |
-
|
| 17 |
-
def set_flag(self, flag_name, value):
|
| 18 |
-
setattr(self, flag_name, value)
|
| 19 |
-
for m in self.children():
|
| 20 |
-
if hasattr(m,'set_flag'):
|
| 21 |
-
m.set_flag(flag_name, value)
|
| 22 |
-
|
| 23 |
-
def forward(self, x):
|
| 24 |
-
for module in self.children():
|
| 25 |
-
x = module(x)
|
| 26 |
-
|
| 27 |
-
kl = 0.0
|
| 28 |
-
for module in self.modules():
|
| 29 |
-
if hasattr(module, 'kl_loss'):
|
| 30 |
-
kl = kl + module.kl_loss()
|
| 31 |
-
|
| 32 |
-
return x, kl
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
class FlattenLayer(ModuleWrapper):
|
| 36 |
-
|
| 37 |
-
def __init__(self, num_features):
|
| 38 |
-
super(FlattenLayer, self).__init__()
|
| 39 |
-
self.num_features = num_features
|
| 40 |
-
|
| 41 |
-
def forward(self, x):
|
| 42 |
-
return x.view(-1, self.num_features)
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
class BBBLinear(ModuleWrapper):
|
| 46 |
-
|
| 47 |
-
def __init__(self, in_features, out_features, bias=True, priors=None):
|
| 48 |
-
super(BBBLinear, self).__init__()
|
| 49 |
-
self.in_features = in_features
|
| 50 |
-
self.out_features = out_features
|
| 51 |
-
self.use_bias = bias
|
| 52 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 53 |
-
|
| 54 |
-
if priors is None:
|
| 55 |
-
priors = {
|
| 56 |
-
'prior_mu': 0,
|
| 57 |
-
'prior_sigma': 0.1,
|
| 58 |
-
'posterior_mu_initial': (0, 0.1),
|
| 59 |
-
'posterior_rho_initial': (-3, 0.1),
|
| 60 |
-
}
|
| 61 |
-
self.prior_mu = priors['prior_mu']
|
| 62 |
-
self.prior_sigma = priors['prior_sigma']
|
| 63 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 64 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 65 |
-
|
| 66 |
-
self.W_mu = Parameter(torch.Tensor(out_features, in_features))
|
| 67 |
-
self.W_rho = Parameter(torch.Tensor(out_features, in_features))
|
| 68 |
-
if self.use_bias:
|
| 69 |
-
self.bias_mu = Parameter(torch.Tensor(out_features))
|
| 70 |
-
self.bias_rho = Parameter(torch.Tensor(out_features))
|
| 71 |
-
else:
|
| 72 |
-
self.register_parameter('bias_mu', None)
|
| 73 |
-
self.register_parameter('bias_rho', None)
|
| 74 |
-
|
| 75 |
-
self.reset_parameters()
|
| 76 |
-
|
| 77 |
-
def reset_parameters(self):
|
| 78 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 79 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 80 |
-
|
| 81 |
-
if self.use_bias:
|
| 82 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 83 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 84 |
-
|
| 85 |
-
def forward(self, x, sample=True):
|
| 86 |
-
|
| 87 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 88 |
-
if self.use_bias:
|
| 89 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 90 |
-
bias_var = self.bias_sigma ** 2
|
| 91 |
-
else:
|
| 92 |
-
self.bias_sigma = bias_var = None
|
| 93 |
-
|
| 94 |
-
act_mu = F.linear(x, self.W_mu, self.bias_mu)
|
| 95 |
-
act_var = 1e-16 + F.linear(x ** 2, self.W_sigma ** 2, bias_var)
|
| 96 |
-
act_std = torch.sqrt(act_var)
|
| 97 |
-
|
| 98 |
-
if self.training or sample:
|
| 99 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 100 |
-
return act_mu + act_std * eps
|
| 101 |
-
else:
|
| 102 |
-
return act_mu
|
| 103 |
-
|
| 104 |
-
def kl_loss(self):
|
| 105 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 106 |
-
if self.use_bias:
|
| 107 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 108 |
-
return kl
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
class BBBConv2d(ModuleWrapper):
|
| 112 |
-
|
| 113 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
|
| 114 |
-
padding=0, dilation=1, bias=True, priors=None):
|
| 115 |
-
super(BBBConv2d, self).__init__()
|
| 116 |
-
self.in_channels = in_channels
|
| 117 |
-
self.out_channels = out_channels
|
| 118 |
-
self.kernel_size = (kernel_size, kernel_size)
|
| 119 |
-
self.stride = stride
|
| 120 |
-
self.padding = padding
|
| 121 |
-
self.dilation = dilation
|
| 122 |
-
self.groups = 1
|
| 123 |
-
self.use_bias = bias
|
| 124 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 125 |
-
|
| 126 |
-
if priors is None:
|
| 127 |
-
priors = {
|
| 128 |
-
'prior_mu': 0,
|
| 129 |
-
'prior_sigma': 0.1,
|
| 130 |
-
'posterior_mu_initial': (0, 0.1),
|
| 131 |
-
'posterior_rho_initial': (-3, 0.1),
|
| 132 |
-
}
|
| 133 |
-
self.prior_mu = priors['prior_mu']
|
| 134 |
-
self.prior_sigma = priors['prior_sigma']
|
| 135 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 136 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 137 |
-
|
| 138 |
-
self.W_mu = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 139 |
-
self.W_rho = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 140 |
-
if self.use_bias:
|
| 141 |
-
self.bias_mu = Parameter(torch.Tensor(out_channels))
|
| 142 |
-
self.bias_rho = Parameter(torch.Tensor(out_channels))
|
| 143 |
-
else:
|
| 144 |
-
self.register_parameter('bias_mu', None)
|
| 145 |
-
self.register_parameter('bias_rho', None)
|
| 146 |
-
|
| 147 |
-
self.reset_parameters()
|
| 148 |
-
|
| 149 |
-
def reset_parameters(self):
|
| 150 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 151 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 152 |
-
|
| 153 |
-
if self.use_bias:
|
| 154 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 155 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 156 |
-
|
| 157 |
-
def forward(self, x, sample=True):
|
| 158 |
-
|
| 159 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 160 |
-
if self.use_bias:
|
| 161 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 162 |
-
bias_var = self.bias_sigma ** 2
|
| 163 |
-
else:
|
| 164 |
-
self.bias_sigma = bias_var = None
|
| 165 |
-
|
| 166 |
-
act_mu = F.conv2d(
|
| 167 |
-
x, self.W_mu, self.bias_mu, self.stride, self.padding, self.dilation, self.groups)
|
| 168 |
-
act_var = 1e-16 + F.conv2d(
|
| 169 |
-
x ** 2, self.W_sigma ** 2, bias_var, self.stride, self.padding, self.dilation, self.groups)
|
| 170 |
-
act_std = torch.sqrt(act_var)
|
| 171 |
-
|
| 172 |
-
if self.training or sample:
|
| 173 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 174 |
-
return act_mu + act_std * eps
|
| 175 |
-
else:
|
| 176 |
-
return act_mu
|
| 177 |
-
|
| 178 |
-
def kl_loss(self):
|
| 179 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 180 |
-
if self.use_bias:
|
| 181 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 182 |
-
return kl
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
def supported_hyperparameters():
|
| 186 |
-
return {'lr','momentum'}
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
class Net(ModuleWrapper):
|
| 190 |
-
|
| 191 |
-
def train_setup(self, prm):
|
| 192 |
-
self.to(self.device)
|
| 193 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 194 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 195 |
-
|
| 196 |
-
def learn(self, train_data):
|
| 197 |
-
for inputs, labels in train_data:
|
| 198 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 199 |
-
self.optimizer.zero_grad()
|
| 200 |
-
outputs = self(inputs)
|
| 201 |
-
loss = self.criteria[0](outputs, labels)
|
| 202 |
-
loss.backward()
|
| 203 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 204 |
-
self.optimizer.step()
|
| 205 |
-
|
| 206 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 207 |
-
super(Net, self).__init__()
|
| 208 |
-
self.device = device
|
| 209 |
-
inputs: int = in_shape[1]
|
| 210 |
-
outputs: int = out_shape[0]
|
| 211 |
-
self.num_classes = outputs
|
| 212 |
-
self.priors = self.priors = {
|
| 213 |
-
'prior_mu': 0,
|
| 214 |
-
'prior_sigma': 0.1,
|
| 215 |
-
'posterior_mu_initial': (0, 0.1),
|
| 216 |
-
'posterior_rho_initial': (-5, 0.1),
|
| 217 |
-
}
|
| 218 |
-
self.act = nn.Softplus
|
| 219 |
-
|
| 220 |
-
self.conv1 = BBBConv2d(inputs, 64, 11, stride=4, padding=5, bias=True, priors=self.priors)
|
| 221 |
-
self.act1 = self.act()
|
| 222 |
-
self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 223 |
-
|
| 224 |
-
self.conv2 = BBBConv2d(64, 192, 5, padding=2, bias=True, priors=self.priors)
|
| 225 |
-
self.act2 = self.act()
|
| 226 |
-
self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 227 |
-
|
| 228 |
-
self.conv3 = BBBConv2d(192, 384, 3, padding=1, bias=True, priors=self.priors)
|
| 229 |
-
self.act3 = self.act()
|
| 230 |
-
|
| 231 |
-
self.conv4 = BBBConv2d(384, 256, 3, padding=1, bias=True, priors=self.priors)
|
| 232 |
-
self.act4 = self.act()
|
| 233 |
-
|
| 234 |
-
self.conv5 = BBBConv2d(256, 128, 3, padding=1, bias=True, priors=self.priors)
|
| 235 |
-
self.act5 = self.act()
|
| 236 |
-
self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 237 |
-
|
| 238 |
-
self.flatten = FlattenLayer(1 * 1 * 128)
|
| 239 |
-
self.classifier = BBBLinear(1 * 1 * 128, outputs, bias=True, priors=self.priors)
|
| 240 |
-
|
| 241 |
-
def __call__(self, *args):
|
| 242 |
-
return super().__call__(*args)[0]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BayesianNet-1.py
DELETED
|
@@ -1,241 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
from torch.nn import Parameter
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def calculate_kl(mu_q, sig_q, mu_p, sig_p):
|
| 8 |
-
kl = 0.5 * (2 * torch.log(sig_p / sig_q) - 1 + (sig_q / sig_p).pow(2) + ((mu_p - mu_q) / sig_p).pow(2)).sum()
|
| 9 |
-
return kl
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
class ModuleWrapper(nn.Module):
|
| 13 |
-
def __init__(self):
|
| 14 |
-
super(ModuleWrapper, self).__init__()
|
| 15 |
-
|
| 16 |
-
def set_flag(self, flag_name, value):
|
| 17 |
-
setattr(self, flag_name, value)
|
| 18 |
-
for m in self.children():
|
| 19 |
-
if hasattr(m, 'set_flag'):
|
| 20 |
-
m.set_flag(flag_name, value)
|
| 21 |
-
|
| 22 |
-
def forward(self, x):
|
| 23 |
-
for module in self.children():
|
| 24 |
-
x = module(x)
|
| 25 |
-
|
| 26 |
-
kl = 0.0
|
| 27 |
-
for module in self.modules():
|
| 28 |
-
if hasattr(module, 'kl_loss'):
|
| 29 |
-
kl = kl + module.kl_loss()
|
| 30 |
-
|
| 31 |
-
return x, kl
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
class FlattenLayer(ModuleWrapper):
|
| 35 |
-
|
| 36 |
-
def __init__(self, num_features):
|
| 37 |
-
super(FlattenLayer, self).__init__()
|
| 38 |
-
self.num_features = num_features
|
| 39 |
-
|
| 40 |
-
def forward(self, x):
|
| 41 |
-
return x.view(-1, self.num_features)
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
class BBBLinear(ModuleWrapper):
|
| 45 |
-
|
| 46 |
-
def __init__(self, in_features, out_features, bias=True, priors=None):
|
| 47 |
-
super(BBBLinear, self).__init__()
|
| 48 |
-
self.in_features = in_features
|
| 49 |
-
self.out_features = out_features
|
| 50 |
-
self.use_bias = bias
|
| 51 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 52 |
-
|
| 53 |
-
if priors is None:
|
| 54 |
-
priors = {
|
| 55 |
-
'prior_mu': 0,
|
| 56 |
-
'prior_sigma': 0.1,
|
| 57 |
-
'posterior_mu_initial': (0, 0.1),
|
| 58 |
-
'posterior_rho_initial': (-3, 0.1),
|
| 59 |
-
}
|
| 60 |
-
self.prior_mu = priors['prior_mu']
|
| 61 |
-
self.prior_sigma = priors['prior_sigma']
|
| 62 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 63 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 64 |
-
|
| 65 |
-
self.W_mu = Parameter(torch.Tensor(out_features, in_features))
|
| 66 |
-
self.W_rho = Parameter(torch.Tensor(out_features, in_features))
|
| 67 |
-
if self.use_bias:
|
| 68 |
-
self.bias_mu = Parameter(torch.Tensor(out_features))
|
| 69 |
-
self.bias_rho = Parameter(torch.Tensor(out_features))
|
| 70 |
-
else:
|
| 71 |
-
self.register_parameter('bias_mu', None)
|
| 72 |
-
self.register_parameter('bias_rho', None)
|
| 73 |
-
|
| 74 |
-
self.reset_parameters()
|
| 75 |
-
|
| 76 |
-
def reset_parameters(self):
|
| 77 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 78 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 79 |
-
|
| 80 |
-
if self.use_bias:
|
| 81 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 82 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 83 |
-
|
| 84 |
-
def forward(self, x, sample=True):
|
| 85 |
-
|
| 86 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 87 |
-
if self.use_bias:
|
| 88 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 89 |
-
bias_var = self.bias_sigma ** 2
|
| 90 |
-
else:
|
| 91 |
-
self.bias_sigma = bias_var = None
|
| 92 |
-
|
| 93 |
-
act_mu = F.linear(x, self.W_mu, self.bias_mu)
|
| 94 |
-
act_var = 1e-16 + F.linear(x ** 2, self.W_sigma ** 2, bias_var)
|
| 95 |
-
act_std = torch.sqrt(act_var)
|
| 96 |
-
|
| 97 |
-
if self.training or sample:
|
| 98 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 99 |
-
return act_mu + act_std * eps
|
| 100 |
-
else:
|
| 101 |
-
return act_mu
|
| 102 |
-
|
| 103 |
-
def kl_loss(self):
|
| 104 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 105 |
-
if self.use_bias:
|
| 106 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 107 |
-
return kl
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
class BBBConv2d(ModuleWrapper):
|
| 111 |
-
|
| 112 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
|
| 113 |
-
padding=0, dilation=1, bias=True, priors=None):
|
| 114 |
-
super(BBBConv2d, self).__init__()
|
| 115 |
-
self.in_channels = in_channels
|
| 116 |
-
self.out_channels = out_channels
|
| 117 |
-
self.kernel_size = (kernel_size, kernel_size)
|
| 118 |
-
self.stride = stride
|
| 119 |
-
self.padding = padding
|
| 120 |
-
self.dilation = dilation
|
| 121 |
-
self.groups = 1
|
| 122 |
-
self.use_bias = bias
|
| 123 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 124 |
-
|
| 125 |
-
if priors is None:
|
| 126 |
-
priors = {
|
| 127 |
-
'prior_mu': 0,
|
| 128 |
-
'prior_sigma': 0.1,
|
| 129 |
-
'posterior_mu_initial': (0, 0.1),
|
| 130 |
-
'posterior_rho_initial': (-3, 0.1),
|
| 131 |
-
}
|
| 132 |
-
self.prior_mu = priors['prior_mu']
|
| 133 |
-
self.prior_sigma = priors['prior_sigma']
|
| 134 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 135 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 136 |
-
|
| 137 |
-
self.W_mu = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 138 |
-
self.W_rho = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 139 |
-
if self.use_bias:
|
| 140 |
-
self.bias_mu = Parameter(torch.Tensor(out_channels))
|
| 141 |
-
self.bias_rho = Parameter(torch.Tensor(out_channels))
|
| 142 |
-
else:
|
| 143 |
-
self.register_parameter('bias_mu', None)
|
| 144 |
-
self.register_parameter('bias_rho', None)
|
| 145 |
-
|
| 146 |
-
self.reset_parameters()
|
| 147 |
-
|
| 148 |
-
def reset_parameters(self):
|
| 149 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 150 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 151 |
-
|
| 152 |
-
if self.use_bias:
|
| 153 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 154 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 155 |
-
|
| 156 |
-
def forward(self, x, sample=True):
|
| 157 |
-
|
| 158 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 159 |
-
if self.use_bias:
|
| 160 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 161 |
-
bias_var = self.bias_sigma ** 2
|
| 162 |
-
else:
|
| 163 |
-
self.bias_sigma = bias_var = None
|
| 164 |
-
|
| 165 |
-
act_mu = F.conv2d(
|
| 166 |
-
x, self.W_mu, self.bias_mu, self.stride, self.padding, self.dilation, self.groups)
|
| 167 |
-
act_var = 1e-16 + F.conv2d(
|
| 168 |
-
x ** 2, self.W_sigma ** 2, bias_var, self.stride, self.padding, self.dilation, self.groups)
|
| 169 |
-
act_std = torch.sqrt(act_var)
|
| 170 |
-
|
| 171 |
-
if self.training or sample:
|
| 172 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 173 |
-
return act_mu + act_std * eps
|
| 174 |
-
else:
|
| 175 |
-
return act_mu
|
| 176 |
-
|
| 177 |
-
def kl_loss(self):
|
| 178 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 179 |
-
if self.use_bias:
|
| 180 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 181 |
-
return kl
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
def supported_hyperparameters():
|
| 185 |
-
return {'lr', 'momentum'}
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
class Net(ModuleWrapper):
|
| 189 |
-
|
| 190 |
-
def train_setup(self, prm):
|
| 191 |
-
self.to(self.device)
|
| 192 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 193 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 194 |
-
|
| 195 |
-
def learn(self, train_data):
|
| 196 |
-
for inputs, labels in train_data:
|
| 197 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 198 |
-
self.optimizer.zero_grad()
|
| 199 |
-
outputs = self(inputs)
|
| 200 |
-
loss = self.criteria[0](outputs, labels)
|
| 201 |
-
loss.backward()
|
| 202 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 203 |
-
self.optimizer.step()
|
| 204 |
-
|
| 205 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 206 |
-
super(Net, self).__init__()
|
| 207 |
-
self.device = device
|
| 208 |
-
inputs: int = in_shape[1]
|
| 209 |
-
outputs: int = out_shape[0]
|
| 210 |
-
self.num_classes = outputs
|
| 211 |
-
self.priors = {
|
| 212 |
-
'prior_mu': 0,
|
| 213 |
-
'prior_sigma': 0.1,
|
| 214 |
-
'posterior_mu_initial': (0, 0.1),
|
| 215 |
-
'posterior_rho_initial': (-5, 0.1),
|
| 216 |
-
}
|
| 217 |
-
self.act = nn.Softplus
|
| 218 |
-
|
| 219 |
-
self.conv1 = BBBConv2d(inputs, 32, 5, padding=2, bias=True, priors=self.priors)
|
| 220 |
-
self.act1 = self.act()
|
| 221 |
-
self.pool1 = nn.MaxPool2d(kernel_size=3, stride=2)
|
| 222 |
-
|
| 223 |
-
self.conv2 = BBBConv2d(32, 64, 5, padding=2, bias=True, priors=self.priors)
|
| 224 |
-
self.act2 = self.act()
|
| 225 |
-
self.pool2 = nn.MaxPool2d(kernel_size=3, stride=2)
|
| 226 |
-
|
| 227 |
-
self.conv3 = BBBConv2d(64, 128, 5, padding=1, bias=True, priors=self.priors)
|
| 228 |
-
self.act3 = self.act()
|
| 229 |
-
self.pool3 = nn.MaxPool2d(kernel_size=3, stride=2)
|
| 230 |
-
|
| 231 |
-
self.flatten = FlattenLayer(2 * 2 * 128)
|
| 232 |
-
self.fc1 = BBBLinear(2 * 2 * 128, 1000, bias=True, priors=self.priors)
|
| 233 |
-
self.act4 = self.act()
|
| 234 |
-
|
| 235 |
-
self.fc2 = BBBLinear(1000, 1000, bias=True, priors=self.priors)
|
| 236 |
-
self.act5 = self.act()
|
| 237 |
-
|
| 238 |
-
self.fc3 = BBBLinear(1000, outputs, bias=True, priors=self.priors)
|
| 239 |
-
|
| 240 |
-
def __call__(self, *args):
|
| 241 |
-
return super().__call__(*args)[0]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/BayesianNet-4f11c8da-cfe1-46ba-b5d0-b5d899929a2e.py
DELETED
|
@@ -1,238 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
from torch.nn import Parameter
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
def calculate_kl(mu_q, sig_q, mu_p, sig_p):
|
| 9 |
-
kl = 0.5 * (2 * torch.log(sig_p / sig_q) - 1 + (sig_q / sig_p).pow(2) + ((mu_p - mu_q) / sig_p).pow(2)).sum()
|
| 10 |
-
return kl
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
class ModuleWrapper(nn.Module):
|
| 14 |
-
def __init__(self):
|
| 15 |
-
super(ModuleWrapper, self).__init__()
|
| 16 |
-
|
| 17 |
-
def set_flag(self, flag_name, value):
|
| 18 |
-
setattr(self, flag_name, value)
|
| 19 |
-
for m in self.children():
|
| 20 |
-
if hasattr(m, 'set_flag'):
|
| 21 |
-
m.set_flag(flag_name, value)
|
| 22 |
-
|
| 23 |
-
def forward(self, x):
|
| 24 |
-
for module in self.children():
|
| 25 |
-
x = module(x)
|
| 26 |
-
|
| 27 |
-
kl = 0.0
|
| 28 |
-
for module in self.modules():
|
| 29 |
-
if hasattr(module, 'kl_loss'):
|
| 30 |
-
kl = kl + module.kl_loss()
|
| 31 |
-
|
| 32 |
-
return x, kl
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
class FlattenLayer(ModuleWrapper):
|
| 36 |
-
|
| 37 |
-
def __init__(self, num_features):
|
| 38 |
-
super(FlattenLayer, self).__init__()
|
| 39 |
-
self.num_features = num_features
|
| 40 |
-
|
| 41 |
-
def forward(self, x):
|
| 42 |
-
return x.view(-1, self.num_features)
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
class BBBLinear(ModuleWrapper):
|
| 46 |
-
|
| 47 |
-
def __init__(self, in_features, out_features, bias=True, priors=None):
|
| 48 |
-
super(BBBLinear, self).__init__()
|
| 49 |
-
self.in_features = in_features
|
| 50 |
-
self.out_features = out_features
|
| 51 |
-
self.use_bias = bias
|
| 52 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 53 |
-
|
| 54 |
-
if priors is None:
|
| 55 |
-
priors = {
|
| 56 |
-
'prior_mu': 0,
|
| 57 |
-
'prior_sigma': 0.1,
|
| 58 |
-
'posterior_mu_initial': (0, 0.1),
|
| 59 |
-
'posterior_rho_initial': (-4, 0.1), # Changed from -3 to -4
|
| 60 |
-
}
|
| 61 |
-
self.prior_mu = priors['prior_mu']
|
| 62 |
-
self.prior_sigma = priors['prior_sigma']
|
| 63 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 64 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 65 |
-
|
| 66 |
-
self.W_mu = Parameter(torch.Tensor(out_features, in_features))
|
| 67 |
-
self.W_rho = Parameter(torch.Tensor(out_features, in_features))
|
| 68 |
-
if self.use_bias:
|
| 69 |
-
self.bias_mu = Parameter(torch.Tensor(out_features))
|
| 70 |
-
self.bias_rho = Parameter(torch.Tensor(out_features))
|
| 71 |
-
else:
|
| 72 |
-
self.register_parameter('bias_mu', None)
|
| 73 |
-
self.register_parameter('bias_rho', None)
|
| 74 |
-
|
| 75 |
-
self.reset_parameters()
|
| 76 |
-
|
| 77 |
-
def reset_parameters(self):
|
| 78 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 79 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 80 |
-
|
| 81 |
-
if self.use_bias:
|
| 82 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 83 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 84 |
-
|
| 85 |
-
def forward(self, x, sample=True):
|
| 86 |
-
|
| 87 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 88 |
-
if self.use_bias:
|
| 89 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 90 |
-
bias_var = self.bias_sigma ** 2
|
| 91 |
-
else:
|
| 92 |
-
self.bias_sigma = bias_var = None
|
| 93 |
-
|
| 94 |
-
act_mu = F.linear(x, self.W_mu, self.bias_mu)
|
| 95 |
-
act_var = 1e-16 + F.linear(x ** 2, self.W_sigma ** 2, bias_var)
|
| 96 |
-
act_std = torch.sqrt(act_var)
|
| 97 |
-
|
| 98 |
-
if self.training or sample:
|
| 99 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 100 |
-
return act_mu + act_std * eps
|
| 101 |
-
else:
|
| 102 |
-
return act_mu
|
| 103 |
-
|
| 104 |
-
def kl_loss(self):
|
| 105 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 106 |
-
if self.use_bias:
|
| 107 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 108 |
-
return kl
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
class BBBConv2d(ModuleWrapper):
|
| 112 |
-
|
| 113 |
-
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
|
| 114 |
-
padding=0, dilation=1, bias=True, priors=None):
|
| 115 |
-
super(BBBConv2d, self).__init__()
|
| 116 |
-
self.in_channels = in_channels
|
| 117 |
-
self.out_channels = out_channels
|
| 118 |
-
self.kernel_size = (kernel_size, kernel_size)
|
| 119 |
-
self.stride = stride
|
| 120 |
-
self.padding = padding
|
| 121 |
-
self.dilation = dilation
|
| 122 |
-
self.groups = 1
|
| 123 |
-
self.use_bias = bias
|
| 124 |
-
self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 125 |
-
|
| 126 |
-
if priors is None:
|
| 127 |
-
priors = {
|
| 128 |
-
'prior_mu': 0,
|
| 129 |
-
'prior_sigma': 0.1,
|
| 130 |
-
'posterior_mu_initial': (0, 0.1),
|
| 131 |
-
'posterior_rho_initial': (-3, 0.1), # Changed from -5 to -6
|
| 132 |
-
}
|
| 133 |
-
self.prior_mu = priors['prior_mu']
|
| 134 |
-
self.prior_sigma = priors['prior_sigma']
|
| 135 |
-
self.posterior_mu_initial = priors['posterior_mu_initial']
|
| 136 |
-
self.posterior_rho_initial = priors['posterior_rho_initial']
|
| 137 |
-
|
| 138 |
-
self.W_mu = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 139 |
-
self.W_rho = Parameter(torch.Tensor(out_channels, in_channels, *self.kernel_size))
|
| 140 |
-
if self.use_bias:
|
| 141 |
-
self.bias_mu = Parameter(torch.Tensor(out_channels))
|
| 142 |
-
self.bias_rho = Parameter(torch.Tensor(out_channels))
|
| 143 |
-
else:
|
| 144 |
-
self.register_parameter('bias_mu', None)
|
| 145 |
-
self.register_parameter('bias_rho', None)
|
| 146 |
-
|
| 147 |
-
self.reset_parameters()
|
| 148 |
-
|
| 149 |
-
def reset_parameters(self):
|
| 150 |
-
self.W_mu.data.normal_(*self.posterior_mu_initial)
|
| 151 |
-
self.W_rho.data.normal_(*self.posterior_rho_initial)
|
| 152 |
-
|
| 153 |
-
if self.use_bias:
|
| 154 |
-
self.bias_mu.data.normal_(*self.posterior_mu_initial)
|
| 155 |
-
self.bias_rho.data.normal_(*self.posterior_rho_initial)
|
| 156 |
-
|
| 157 |
-
def forward(self, x, sample=True):
|
| 158 |
-
|
| 159 |
-
self.W_sigma = torch.log1p(torch.exp(self.W_rho))
|
| 160 |
-
if self.use_bias:
|
| 161 |
-
self.bias_sigma = torch.log1p(torch.exp(self.bias_rho))
|
| 162 |
-
bias_var = self.bias_sigma ** 2
|
| 163 |
-
else:
|
| 164 |
-
self.bias_sigma = bias_var = None
|
| 165 |
-
|
| 166 |
-
act_mu = F.conv2d(
|
| 167 |
-
x, self.W_mu, self.bias_mu, self.stride, self.padding, self.dilation, self.groups)
|
| 168 |
-
act_var = 1e-16 + F.conv2d(
|
| 169 |
-
x ** 2, self.W_sigma ** 2, bias_var, self.stride, self.padding, self.dilation, self.groups)
|
| 170 |
-
act_std = torch.sqrt(act_var)
|
| 171 |
-
|
| 172 |
-
if self.training or sample:
|
| 173 |
-
eps = torch.empty(act_mu.size()).normal_(0, 1).to(self.device)
|
| 174 |
-
return act_mu + act_std * eps
|
| 175 |
-
else:
|
| 176 |
-
return act_mu
|
| 177 |
-
|
| 178 |
-
def kl_loss(self):
|
| 179 |
-
kl = calculate_kl(self.prior_mu, self.prior_sigma, self.W_mu, self.W_sigma)
|
| 180 |
-
if self.use_bias:
|
| 181 |
-
kl += calculate_kl(self.prior_mu, self.prior_sigma, self.bias_mu, self.bias_sigma)
|
| 182 |
-
return kl
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
def supported_hyperparameters():
|
| 186 |
-
return {'lr', 'momentum'}
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
class Net(ModuleWrapper):
|
| 190 |
-
|
| 191 |
-
def train_setup(self, prm):
|
| 192 |
-
self.to(self.device)
|
| 193 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 194 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 195 |
-
|
| 196 |
-
def learn(self, train_data):
|
| 197 |
-
for inputs, labels in train_data:
|
| 198 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 199 |
-
self.optimizer.zero_grad()
|
| 200 |
-
outputs = self(inputs)
|
| 201 |
-
loss = self.criteria[0](outputs, labels)
|
| 202 |
-
loss.backward()
|
| 203 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 204 |
-
self.optimizer.step()
|
| 205 |
-
|
| 206 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 207 |
-
super(Net, self).__init__()
|
| 208 |
-
self.device = device
|
| 209 |
-
inputs: int = in_shape[1]
|
| 210 |
-
outputs: int = out_shape[0]
|
| 211 |
-
self.num_classes = outputs
|
| 212 |
-
self.priors = self.priors = {
|
| 213 |
-
'prior_mu': 0,
|
| 214 |
-
'prior_sigma': 0.1,
|
| 215 |
-
'posterior_mu_initial': (0, 0.1),
|
| 216 |
-
'posterior_rho_initial': (-4, 0.1), # Changed from -5 to -6
|
| 217 |
-
}
|
| 218 |
-
self.act = nn.Softplus
|
| 219 |
-
|
| 220 |
-
self.conv1 = BBBConv2d(inputs, 6, 5, padding=0, bias=True, priors=self.priors)
|
| 221 |
-
self.act1 = self.act()
|
| 222 |
-
self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 223 |
-
|
| 224 |
-
self.conv2 = BBBConv2d(6, 16, 5, padding=0, bias=True, priors=self.priors)
|
| 225 |
-
self.act2 = self.act()
|
| 226 |
-
self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
|
| 227 |
-
|
| 228 |
-
self.flatten = FlattenLayer(5 * 5 * 16)
|
| 229 |
-
self.fc1 = BBBLinear(5 * 5 * 16, 120, bias=True, priors=self.priors)
|
| 230 |
-
self.act3 = self.act()
|
| 231 |
-
|
| 232 |
-
self.fc2 = BBBLinear(120, 84, bias=True, priors=self.priors)
|
| 233 |
-
self.act4 = self.act()
|
| 234 |
-
|
| 235 |
-
self.fc3 = BBBLinear(84, outputs, bias=True, priors=self.priors)
|
| 236 |
-
|
| 237 |
-
def __call__(self, *args):
|
| 238 |
-
return super().__call__(*args)[0]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C10C-RESNETLSTM-6a517327bf0ef897a22186a2061e85b3.py
DELETED
|
@@ -1,180 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
|
| 5 |
-
def supported_hyperparameters():
|
| 6 |
-
return {'lr','momentum'}
|
| 7 |
-
|
| 8 |
-
class SEBlock(nn.Module):
|
| 9 |
-
def __init__(self, channel, reduction=4):
|
| 10 |
-
super().__init__()
|
| 11 |
-
self.avg = nn.AdaptiveAvgPool2d(1)
|
| 12 |
-
self.fc1 = nn.Linear(channel, channel // reduction, bias=False)
|
| 13 |
-
self.fc2 = nn.Linear(channel // reduction, channel, bias=False)
|
| 14 |
-
|
| 15 |
-
def forward(self, x):
|
| 16 |
-
b, c, _, _ = x.size()
|
| 17 |
-
y = self.avg(x).view(b, c)
|
| 18 |
-
y = self.fc2(F.relu(self.fc1(y), inplace=True)).view(b, c, 1, 1)
|
| 19 |
-
return x * torch.sigmoid(y)
|
| 20 |
-
|
| 21 |
-
class DepthwiseSeparableConv(nn.Module):
|
| 22 |
-
def __init__(self, in_ch, out_ch, k=3, s=1, p=1):
|
| 23 |
-
super().__init__()
|
| 24 |
-
self.dw = nn.Conv2d(in_ch, in_ch, k, s, p, groups=in_ch, bias=False)
|
| 25 |
-
self.pw = nn.Conv2d(in_ch, out_ch, 1, 1, bias=False)
|
| 26 |
-
self.bn = nn.BatchNorm2d(out_ch)
|
| 27 |
-
|
| 28 |
-
def forward(self, x):
|
| 29 |
-
x = self.dw(x)
|
| 30 |
-
x = self.pw(x)
|
| 31 |
-
return F.relu(self.bn(x), inplace=True)
|
| 32 |
-
|
| 33 |
-
class YourEncoder(nn.Module):
|
| 34 |
-
def __init__(self, in_channels, hidden_dim=512):
|
| 35 |
-
super().__init__()
|
| 36 |
-
h2 = hidden_dim // 2
|
| 37 |
-
self.stem = nn.Sequential(
|
| 38 |
-
nn.Conv2d(in_channels, h2, 3, 2, 1, bias=False),
|
| 39 |
-
nn.BatchNorm2d(h2),
|
| 40 |
-
nn.ReLU(inplace=True),
|
| 41 |
-
DepthwiseSeparableConv(h2, hidden_dim),
|
| 42 |
-
SEBlock(hidden_dim),
|
| 43 |
-
nn.AdaptiveAvgPool2d((1,1))
|
| 44 |
-
)
|
| 45 |
-
self.fc = nn.Linear(hidden_dim, hidden_dim)
|
| 46 |
-
|
| 47 |
-
def forward(self, x):
|
| 48 |
-
x = self.stem(x)
|
| 49 |
-
x = x.view(x.size(0), -1)
|
| 50 |
-
x = self.fc(x)
|
| 51 |
-
return x
|
| 52 |
-
|
| 53 |
-
class Attention(nn.Module):
|
| 54 |
-
def __init__(self, hidden_size, feature_dim):
|
| 55 |
-
super().__init__()
|
| 56 |
-
self.q = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 57 |
-
self.k = nn.Linear(feature_dim, hidden_size, bias=False)
|
| 58 |
-
self.v = nn.Linear(feature_dim, hidden_size, bias=False)
|
| 59 |
-
|
| 60 |
-
def forward(self, h, feats):
|
| 61 |
-
if feats.dim()==2:
|
| 62 |
-
feats = feats.unsqueeze(1)
|
| 63 |
-
q = self.q(h)
|
| 64 |
-
k = self.k(feats)
|
| 65 |
-
v = self.v(feats)
|
| 66 |
-
score = torch.einsum('bh,brh->br', q, k)
|
| 67 |
-
attn = F.softmax(score, dim=1)
|
| 68 |
-
ctx = torch.einsum('br,brh->bh', attn, v)
|
| 69 |
-
return ctx
|
| 70 |
-
|
| 71 |
-
class YourDecoder(nn.Module):
|
| 72 |
-
def __init__(self, vocab_size, feature_dim=512, hidden_size=512):
|
| 73 |
-
super().__init__()
|
| 74 |
-
self.embed = nn.Embedding(vocab_size, hidden_size)
|
| 75 |
-
self.attn = Attention(hidden_size, feature_dim)
|
| 76 |
-
self.cell = nn.GRUCell(input_size=hidden_size*2, hidden_size=hidden_size)
|
| 77 |
-
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 78 |
-
self.hidden_size = hidden_size
|
| 79 |
-
self.vocab_size = vocab_size
|
| 80 |
-
|
| 81 |
-
def init_zero_hidden(self, batch, device):
|
| 82 |
-
h0 = torch.zeros(batch, self.hidden_size, device=device)
|
| 83 |
-
c0 = torch.zeros(batch, self.hidden_size, device=device)
|
| 84 |
-
return (h0, c0)
|
| 85 |
-
|
| 86 |
-
def forward(self, inputs, hidden_state, features):
|
| 87 |
-
B, T = inputs.size()
|
| 88 |
-
if features.dim()==3 and features.size(1)==1:
|
| 89 |
-
features = features.squeeze(1)
|
| 90 |
-
if hidden_state is None or hidden_state[0].size(0)!=B:
|
| 91 |
-
h = torch.zeros(B, self.hidden_size, device=inputs.device)
|
| 92 |
-
else:
|
| 93 |
-
h = hidden_state[0]
|
| 94 |
-
embs = self.embed(inputs)
|
| 95 |
-
outs = []
|
| 96 |
-
for t in range(T):
|
| 97 |
-
ctx = self.attn(h, features)
|
| 98 |
-
x = torch.cat([embs[:, t, :], ctx], dim=1)
|
| 99 |
-
h = self.cell(x, h)
|
| 100 |
-
outs.append(self.fc(h))
|
| 101 |
-
logits = torch.stack(outs, dim=1)
|
| 102 |
-
return logits, (h, torch.zeros_like(h))
|
| 103 |
-
|
| 104 |
-
@torch.no_grad()
|
| 105 |
-
def greedy_decode(self, features, max_len=50, start_id=1, end_id=2):
|
| 106 |
-
if features.dim()==3 and features.size(1)==1:
|
| 107 |
-
features = features.squeeze(1)
|
| 108 |
-
B = features.size(0)
|
| 109 |
-
device = features.device
|
| 110 |
-
h = torch.zeros(B, self.hidden_size, device=device)
|
| 111 |
-
cur = torch.full((B,), start_id, dtype=torch.long, device=device)
|
| 112 |
-
tokens = []
|
| 113 |
-
for _ in range(max_len):
|
| 114 |
-
emb = self.embed(cur)
|
| 115 |
-
ctx = self.attn(h, features)
|
| 116 |
-
x = torch.cat([emb, ctx], dim=1)
|
| 117 |
-
h = self.cell(x, h)
|
| 118 |
-
logit = self.fc(h)
|
| 119 |
-
cur = logit.argmax(dim=1)
|
| 120 |
-
tokens.append(cur)
|
| 121 |
-
if (cur==end_id).all():
|
| 122 |
-
break
|
| 123 |
-
if len(tokens)==0:
|
| 124 |
-
return torch.empty(B, 0, dtype=torch.long, device=device)
|
| 125 |
-
return torch.stack(tokens, dim=1)
|
| 126 |
-
|
| 127 |
-
class Net(nn.Module):
|
| 128 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 129 |
-
super().__init__()
|
| 130 |
-
self.device = device
|
| 131 |
-
in_channels = int(in_shape[1])
|
| 132 |
-
vocab_size = int(out_shape[0])
|
| 133 |
-
hidden = 512
|
| 134 |
-
self.encoder = YourEncoder(in_channels, hidden_dim=hidden)
|
| 135 |
-
self.rnn = YourDecoder(vocab_size, feature_dim=hidden, hidden_size=hidden)
|
| 136 |
-
self.criterion = nn.CrossEntropyLoss(ignore_index=0)
|
| 137 |
-
self.optimizer = None
|
| 138 |
-
self.vocab_size = vocab_size
|
| 139 |
-
|
| 140 |
-
def _norm_caps(self, caps):
|
| 141 |
-
if caps.ndim==3:
|
| 142 |
-
caps = caps[:,0,:]
|
| 143 |
-
elif caps.ndim==1:
|
| 144 |
-
caps = caps.unsqueeze(0)
|
| 145 |
-
return caps.long()
|
| 146 |
-
|
| 147 |
-
def forward(self, images, captions=None, hidden_state=None):
|
| 148 |
-
assert images.dim()==4
|
| 149 |
-
feats = self.encoder(images)
|
| 150 |
-
B = images.size(0)
|
| 151 |
-
if captions is None:
|
| 152 |
-
return self.rnn.greedy_decode(feats, max_len=50)
|
| 153 |
-
caps = self._norm_caps(captions)
|
| 154 |
-
inputs = caps[:, :-1]
|
| 155 |
-
if hidden_state is None or (isinstance(hidden_state, tuple) and hidden_state[0].size(0)!=B):
|
| 156 |
-
hidden_state = self.rnn.init_zero_hidden(B, images.device)
|
| 157 |
-
logits, _ = self.rnn(inputs, hidden_state, feats)
|
| 158 |
-
assert logits.dim()==3 and logits.size(1)==inputs.size(1)
|
| 159 |
-
return logits
|
| 160 |
-
|
| 161 |
-
def train_setup(self, prm):
|
| 162 |
-
self.to(self.device)
|
| 163 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 164 |
-
self.criterion = self.criterion.to(self.device)
|
| 165 |
-
|
| 166 |
-
def learn(self, train_data):
|
| 167 |
-
self.train()
|
| 168 |
-
for images, captions in train_data:
|
| 169 |
-
images = images.to(self.device)
|
| 170 |
-
captions = captions.to(self.device)
|
| 171 |
-
caps = self._norm_caps(captions)
|
| 172 |
-
self.optimizer.zero_grad()
|
| 173 |
-
logits = self(images, caps, None)
|
| 174 |
-
T = min(logits.size(1), caps.size(1)-1)
|
| 175 |
-
logits = logits[:, :T, :]
|
| 176 |
-
tgt = caps[:, 1:1+T]
|
| 177 |
-
loss = self.criterion(logits.reshape(-1, self.vocab_size), tgt.reshape(-1))
|
| 178 |
-
loss.backward()
|
| 179 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 180 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C10C-RESNETLSTM-8f7ac9c241d5b9546f8cd3484e0e100b.py
DELETED
|
@@ -1,245 +0,0 @@
|
|
| 1 |
-
import math
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
|
| 6 |
-
# Optional: discover PAD/BOS/EOS ids from loader
|
| 7 |
-
try:
|
| 8 |
-
from ab.nn.loader.coco_.Caption import GLOBAL_CAPTION_VOCAB
|
| 9 |
-
except Exception:
|
| 10 |
-
GLOBAL_CAPTION_VOCAB = {}
|
| 11 |
-
|
| 12 |
-
def supported_hyperparameters():
|
| 13 |
-
return {'lr', 'momentum', 'dropout'}
|
| 14 |
-
|
| 15 |
-
# ---------- helpers ----------
|
| 16 |
-
def _special_ids(vocab: dict, vocab_size: int):
|
| 17 |
-
def hit(keys, default):
|
| 18 |
-
for k in keys:
|
| 19 |
-
if k in vocab:
|
| 20 |
-
return int(vocab[k])
|
| 21 |
-
return max(0, min(default, vocab_size - 1))
|
| 22 |
-
pad = hit(['<PAD>', '<pad>', '<pad_token>', '<blank>', '<null>'], 0)
|
| 23 |
-
bos = hit(['<BOS>', '<bos>', '<s>', '<start>', '<SOS>', '<sos>'], 1)
|
| 24 |
-
eos = hit(['<EOS>', '<eos>', '</s>', '<end>', '<EOS_TOKEN>'], 2)
|
| 25 |
-
return pad, bos, eos
|
| 26 |
-
|
| 27 |
-
class PositionalEncoding(nn.Module):
|
| 28 |
-
def __init__(self, d_model: int, max_len: int = 4096, dropout: float = 0.0):
|
| 29 |
-
super().__init__()
|
| 30 |
-
pe = torch.zeros(max_len, d_model)
|
| 31 |
-
pos = torch.arange(0, max_len).float().unsqueeze(1)
|
| 32 |
-
div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
|
| 33 |
-
pe[:, 0::2] = torch.sin(pos * div)
|
| 34 |
-
pe[:, 1::2] = torch.cos(pos * div)
|
| 35 |
-
self.register_buffer('pe', pe.unsqueeze(0), persistent=False) # (1, L, D)
|
| 36 |
-
self.drop = nn.Dropout(dropout)
|
| 37 |
-
|
| 38 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 39 |
-
# x: (B, L, D)
|
| 40 |
-
L = x.size(1)
|
| 41 |
-
return self.drop(x + self.pe[:, :L, :])
|
| 42 |
-
|
| 43 |
-
# ---------- encoder ----------
|
| 44 |
-
class BagNetBlock(nn.Module):
|
| 45 |
-
def __init__(self, in_ch, out_ch, k=3, s=1):
|
| 46 |
-
super().__init__()
|
| 47 |
-
mid = max(1, out_ch // 4)
|
| 48 |
-
self.conv1 = nn.Conv2d(in_ch, mid, 1, 1, 0, bias=False)
|
| 49 |
-
self.conv2 = nn.Conv2d(mid, mid, k, s, (k - 1)//2, bias=False)
|
| 50 |
-
self.bn2 = nn.BatchNorm2d(mid)
|
| 51 |
-
self.conv3 = nn.Conv2d(mid, out_ch, 1, 1, 0, bias=False)
|
| 52 |
-
self.proj = None if (in_ch == out_ch and s == 1) else nn.Conv2d(in_ch, out_ch, 1, s, 0, bias=False)
|
| 53 |
-
self.act = nn.ReLU(inplace=True)
|
| 54 |
-
|
| 55 |
-
def forward(self, x):
|
| 56 |
-
idt = x if self.proj is None else self.proj(x)
|
| 57 |
-
y = self.conv1(x)
|
| 58 |
-
y = self.conv2(y); y = self.bn2(y); y = self.act(y)
|
| 59 |
-
y = self.conv3(y)
|
| 60 |
-
if y.shape[-2:] != idt.shape[-2:]:
|
| 61 |
-
idt = F.interpolate(idt, size=y.shape[-2:], mode='bilinear', align_corners=False)
|
| 62 |
-
return self.act(y + idt)
|
| 63 |
-
|
| 64 |
-
class CNNEncoder(nn.Module):
|
| 65 |
-
def __init__(self, in_ch: int, feat_ch: int = 384):
|
| 66 |
-
super().__init__()
|
| 67 |
-
self.stem = nn.Sequential(
|
| 68 |
-
nn.Conv2d(in_ch, 64, 3, 2, 1, bias=False),
|
| 69 |
-
nn.BatchNorm2d(64), nn.SiLU(),
|
| 70 |
-
nn.Conv2d(64, 128, 3, 2, 1, bias=False),
|
| 71 |
-
nn.BatchNorm2d(128), nn.SiLU(),
|
| 72 |
-
)
|
| 73 |
-
self.b1 = BagNetBlock(128, 256, k=3, s=2) # (H/8, W/8)
|
| 74 |
-
self.b2 = BagNetBlock(256, feat_ch, k=3, s=1)
|
| 75 |
-
|
| 76 |
-
def forward(self, x):
|
| 77 |
-
x = self.stem(x)
|
| 78 |
-
x = self.b1(x)
|
| 79 |
-
x = self.b2(x) # (B, C, h, w)
|
| 80 |
-
return x
|
| 81 |
-
|
| 82 |
-
# ---------- decoder ----------
|
| 83 |
-
class TransformerDecoder(nn.Module):
|
| 84 |
-
def __init__(self, vocab_size: int, d_model: int = 512, nhead: int = 8, num_layers: int = 4,
|
| 85 |
-
dropout: float = 0.1, pad_idx: int = 0):
|
| 86 |
-
super().__init__()
|
| 87 |
-
self.pad_idx = pad_idx
|
| 88 |
-
self.embed = nn.Embedding(vocab_size, d_model, padding_idx=pad_idx)
|
| 89 |
-
self.pos = PositionalEncoding(d_model, dropout=dropout)
|
| 90 |
-
layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead,
|
| 91 |
-
dim_feedforward=2048, batch_first=True,
|
| 92 |
-
dropout=dropout, activation='gelu')
|
| 93 |
-
self.dec = nn.TransformerDecoder(layer, num_layers=num_layers)
|
| 94 |
-
self.fc = nn.Linear(d_model, vocab_size)
|
| 95 |
-
|
| 96 |
-
@staticmethod
|
| 97 |
-
def _causal_mask(L, device, dtype):
|
| 98 |
-
# Make the dtype consistent with PyTorch recommendations to avoid warnings
|
| 99 |
-
m = torch.full((L, L), float('-inf'), device=device, dtype=dtype)
|
| 100 |
-
return torch.triu(m, diagonal=1)
|
| 101 |
-
|
| 102 |
-
def forward(self, tgt_tokens: torch.Tensor, memory: torch.Tensor) -> torch.Tensor:
|
| 103 |
-
"""
|
| 104 |
-
tgt_tokens: (B, T), memory: (B, S, D)
|
| 105 |
-
returns logits: (B, T, V)
|
| 106 |
-
"""
|
| 107 |
-
B, T = tgt_tokens.shape
|
| 108 |
-
x = self.embed(tgt_tokens) # (B,T,D)
|
| 109 |
-
x = self.pos(x)
|
| 110 |
-
# Use float mask to match attn_mask dtype
|
| 111 |
-
tgt_mask = self._causal_mask(T, x.device, x.dtype) # (T,T)
|
| 112 |
-
tgt_kpm = (tgt_tokens == self.pad_idx) # (B,T) bool
|
| 113 |
-
y = self.dec(tgt=x, memory=memory,
|
| 114 |
-
tgt_mask=tgt_mask,
|
| 115 |
-
tgt_key_padding_mask=tgt_kpm)
|
| 116 |
-
return self.fc(y) # (B,T,V)
|
| 117 |
-
|
| 118 |
-
# ---------- full model ----------
|
| 119 |
-
class Net(nn.Module):
|
| 120 |
-
"""
|
| 121 |
-
Returns a Tensor from forward() (never a tuple), so metrics like BLEU can call .dim().
|
| 122 |
-
API:
|
| 123 |
-
- __init__(in_shape, out_shape, prm, device)
|
| 124 |
-
- forward(images, captions=None) -> logits Tensor
|
| 125 |
-
- train_setup(prm)
|
| 126 |
-
- learn(train_data)
|
| 127 |
-
"""
|
| 128 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 129 |
-
super().__init__()
|
| 130 |
-
self.device = device
|
| 131 |
-
self.in_channels = int(in_shape[1])
|
| 132 |
-
self.vocab_size = int(out_shape[0])
|
| 133 |
-
|
| 134 |
-
# Hyperparams (consumed)
|
| 135 |
-
self.dropout_p = float(prm.get('dropout', 0.1))
|
| 136 |
-
self.max_len = int(prm.get('max_len', 20))
|
| 137 |
-
|
| 138 |
-
# Special tokens
|
| 139 |
-
self.pad_idx, self.bos_id, self.eos_id = _special_ids(GLOBAL_CAPTION_VOCAB or {}, self.vocab_size)
|
| 140 |
-
|
| 141 |
-
# Encoder -> sequence of d_model features
|
| 142 |
-
d_model = 512
|
| 143 |
-
enc_feat_ch = 384
|
| 144 |
-
self.encoder = CNNEncoder(self.in_channels, feat_ch=enc_feat_ch)
|
| 145 |
-
self.enc_proj = nn.Linear(enc_feat_ch, d_model)
|
| 146 |
-
self.enc_pos = PositionalEncoding(d_model, dropout=self.dropout_p)
|
| 147 |
-
self.enc_drop = nn.Dropout(self.dropout_p)
|
| 148 |
-
|
| 149 |
-
# Transformer decoder
|
| 150 |
-
self.decoder = TransformerDecoder(self.vocab_size, d_model=d_model, nhead=8,
|
| 151 |
-
num_layers=4, dropout=self.dropout_p, pad_idx=self.pad_idx)
|
| 152 |
-
|
| 153 |
-
# Training attrs init in train_setup
|
| 154 |
-
self.criteria = None
|
| 155 |
-
self.optimizer = None
|
| 156 |
-
self.scaler = None
|
| 157 |
-
|
| 158 |
-
# -- encoder helper --
|
| 159 |
-
def _encode(self, images: torch.Tensor) -> torch.Tensor:
|
| 160 |
-
f = self.encoder(images) # (B,C,h,w)
|
| 161 |
-
B, C, h, w = f.shape
|
| 162 |
-
seq = f.view(B, C, h*w).permute(0, 2, 1) # (B,S,C)
|
| 163 |
-
seq = self.enc_proj(seq) # (B,S,D)
|
| 164 |
-
seq = self.enc_pos(seq)
|
| 165 |
-
seq = self.enc_drop(seq)
|
| 166 |
-
return seq # (B,S,D)
|
| 167 |
-
|
| 168 |
-
# -- forward --
|
| 169 |
-
def forward(self, images, captions=None):
|
| 170 |
-
"""
|
| 171 |
-
Training (teacher forcing):
|
| 172 |
-
inputs = captions[:, :-1] -> logits over positions 1..T-1
|
| 173 |
-
returns logits: (B, T-1, V)
|
| 174 |
-
Inference (captions=None):
|
| 175 |
-
greedy decode up to max_len
|
| 176 |
-
returns logits: (B, L, V) of generated steps
|
| 177 |
-
"""
|
| 178 |
-
assert images.dim() == 4, "images must be (B,C,H,W)"
|
| 179 |
-
memory = self._encode(images) # (B,S,D)
|
| 180 |
-
|
| 181 |
-
if captions is not None:
|
| 182 |
-
if captions.ndim == 3:
|
| 183 |
-
captions = captions[:, 0, :] # (B,T)
|
| 184 |
-
inputs = captions[:, :-1] # (B,T-1)
|
| 185 |
-
logits = self.decoder(inputs, memory) # (B,T-1,V)
|
| 186 |
-
return logits # Tensor ONLY
|
| 187 |
-
|
| 188 |
-
# Inference: greedy
|
| 189 |
-
B = images.size(0)
|
| 190 |
-
device = images.device
|
| 191 |
-
cur = torch.full((B, 1), self.bos_id, dtype=torch.long, device=device)
|
| 192 |
-
steps = []
|
| 193 |
-
for _ in range(self.max_len):
|
| 194 |
-
step_logits = self.decoder(cur, memory)[:, -1:, :] # (B,1,V)
|
| 195 |
-
steps.append(step_logits)
|
| 196 |
-
next_tok = step_logits.argmax(dim=-1) # (B,1)
|
| 197 |
-
cur = torch.cat([cur, next_tok], dim=1)
|
| 198 |
-
if (next_tok.squeeze(1) == self.eos_id).all():
|
| 199 |
-
break
|
| 200 |
-
logits = torch.cat(steps, dim=1) if steps else torch.zeros((B, 0, self.vocab_size), device=device)
|
| 201 |
-
return logits # Tensor ONLY
|
| 202 |
-
|
| 203 |
-
# -- training setup --
|
| 204 |
-
def train_setup(self, prm):
|
| 205 |
-
self.to(self.device)
|
| 206 |
-
# Loss: ignore PAD; a bit of label smoothing helps BLEU
|
| 207 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=self.pad_idx, label_smoothing=0.1).to(self.device),)
|
| 208 |
-
# Consume 'momentum' by mapping to AdamW beta1
|
| 209 |
-
beta1 = float(prm.get('momentum', 0.9))
|
| 210 |
-
self.optimizer = torch.optim.AdamW(self.parameters(),
|
| 211 |
-
lr=float(prm['lr']),
|
| 212 |
-
betas=(beta1, 0.999),
|
| 213 |
-
weight_decay=1e-4)
|
| 214 |
-
# New AMP API to avoid deprecation warning
|
| 215 |
-
self.scaler = torch.amp.GradScaler('cuda', enabled=(self.device.type == 'cuda'))
|
| 216 |
-
|
| 217 |
-
# -- one epoch training loop --
|
| 218 |
-
def learn(self, train_data):
|
| 219 |
-
"""
|
| 220 |
-
Expects batches like (images, captions, *rest).
|
| 221 |
-
"""
|
| 222 |
-
assert self.criteria and self.optimizer is not None and self.scaler is not None, "Call train_setup(prm) first."
|
| 223 |
-
self.train()
|
| 224 |
-
amp_device = 'cuda' if self.device.type == 'cuda' else 'cpu'
|
| 225 |
-
for batch in train_data:
|
| 226 |
-
if isinstance(batch, (list, tuple)):
|
| 227 |
-
images, captions = batch[0], batch[1]
|
| 228 |
-
else:
|
| 229 |
-
images, captions = batch
|
| 230 |
-
images = images.to(self.device, non_blocking=True)
|
| 231 |
-
captions = captions.to(self.device, non_blocking=True)
|
| 232 |
-
|
| 233 |
-
with torch.amp.autocast(amp_device, enabled=(self.device.type == 'cuda')):
|
| 234 |
-
if captions.ndim == 3:
|
| 235 |
-
captions = captions[:, 0, :]
|
| 236 |
-
logits = self.forward(images, captions) # (B,T-1,V) Tensor
|
| 237 |
-
targets = captions[:, 1:] # (B,T-1)
|
| 238 |
-
loss = self.criteria[0](logits.reshape(-1, logits.size(-1)),
|
| 239 |
-
targets.reshape(-1))
|
| 240 |
-
|
| 241 |
-
self.optimizer.zero_grad(set_to_none=True)
|
| 242 |
-
self.scaler.scale(loss).backward()
|
| 243 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 244 |
-
self.scaler.step(self.optimizer)
|
| 245 |
-
self.scaler.update()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C10C-RESNETLSTM-IMG-CAP-IMPROVED.py
DELETED
|
@@ -1,230 +0,0 @@
|
|
| 1 |
-
import math
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
import torchvision.models as tv
|
| 6 |
-
|
| 7 |
-
# --- AlexNet weights (safe fallback for older torchvision) ---
|
| 8 |
-
try:
|
| 9 |
-
from torchvision.models import AlexNet_Weights
|
| 10 |
-
ALEXNET_W = AlexNet_Weights.IMAGENET1K_V1
|
| 11 |
-
except Exception:
|
| 12 |
-
ALEXNET_W = None # fallback: will use pretrained=True on older torchvision
|
| 13 |
-
|
| 14 |
-
# Optional: discover PAD/BOS/EOS ids from loader
|
| 15 |
-
try:
|
| 16 |
-
from ab.nn.loader.coco_.Caption import GLOBAL_CAPTION_VOCAB
|
| 17 |
-
except Exception:
|
| 18 |
-
GLOBAL_CAPTION_VOCAB = {}
|
| 19 |
-
|
| 20 |
-
def supported_hyperparameters():
|
| 21 |
-
# repo's train.py consumes lr/momentum/dropout via -p JSON or optuna ranges
|
| 22 |
-
return {'lr', 'momentum', 'dropout'}
|
| 23 |
-
|
| 24 |
-
# ---------- helpers ----------
|
| 25 |
-
def _special_ids(vocab: dict, vocab_size: int):
|
| 26 |
-
def hit(keys, default):
|
| 27 |
-
for k in keys:
|
| 28 |
-
if k in vocab:
|
| 29 |
-
return int(vocab[k])
|
| 30 |
-
return max(0, min(default, vocab_size - 1))
|
| 31 |
-
pad = hit(['<PAD>', '<pad>', '<pad_token>', '<blank>', '<null>'], 0)
|
| 32 |
-
bos = hit(['<BOS>', '<bos>', '<s>', '<start>', '<SOS>', '<sos>'], 1)
|
| 33 |
-
eos = hit(['<EOS>', '<eos>', '</s>', '<end>', '<EOS_TOKEN>'], 2)
|
| 34 |
-
return pad, bos, eos
|
| 35 |
-
|
| 36 |
-
class PositionalEncoding(nn.Module):
|
| 37 |
-
def __init__(self, d_model: int, max_len: int = 4096, dropout: float = 0.0):
|
| 38 |
-
super().__init__()
|
| 39 |
-
pe = torch.zeros(max_len, d_model)
|
| 40 |
-
pos = torch.arange(0, max_len).float().unsqueeze(1)
|
| 41 |
-
div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
|
| 42 |
-
pe[:, 0::2] = torch.sin(pos * div)
|
| 43 |
-
pe[:, 1::2] = torch.cos(pos * div)
|
| 44 |
-
self.register_buffer('pe', pe.unsqueeze(0), persistent=False) # (1, L, D)
|
| 45 |
-
self.drop = nn.Dropout(dropout)
|
| 46 |
-
|
| 47 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 48 |
-
# x: (B, L, D)
|
| 49 |
-
L = x.size(1)
|
| 50 |
-
return self.drop(x + self.pe[:, :L, :])
|
| 51 |
-
|
| 52 |
-
# ---------- encoder (AlexNet features) ----------
|
| 53 |
-
class CNNEncoder(nn.Module):
|
| 54 |
-
"""
|
| 55 |
-
AlexNet conv feature extractor.
|
| 56 |
-
Output: (B, 256, h, w)
|
| 57 |
-
"""
|
| 58 |
-
def __init__(self, in_ch: int, feat_ch: int = 256): # feat_ch kept for signature
|
| 59 |
-
super().__init__()
|
| 60 |
-
# torchvision AlexNet expects 3-channel RGB, repo's transforms should handle normalization/resize
|
| 61 |
-
if ALEXNET_W is None:
|
| 62 |
-
self.backbone = tv.alexnet(pretrained=True).features
|
| 63 |
-
else:
|
| 64 |
-
self.backbone = tv.alexnet(weights=ALEXNET_W).features
|
| 65 |
-
|
| 66 |
-
def forward(self, x):
|
| 67 |
-
return self.backbone(x) # (B, 256, h, w)
|
| 68 |
-
|
| 69 |
-
# ---------- decoder ----------
|
| 70 |
-
class TransformerDecoder(nn.Module):
|
| 71 |
-
def __init__(self, vocab_size: int, d_model: int = 512, nhead: int = 8, num_layers: int = 4,
|
| 72 |
-
dropout: float = 0.1, pad_idx: int = 0):
|
| 73 |
-
super().__init__()
|
| 74 |
-
self.pad_idx = pad_idx
|
| 75 |
-
self.embed = nn.Embedding(vocab_size, d_model, padding_idx=pad_idx)
|
| 76 |
-
self.pos = PositionalEncoding(d_model, dropout=dropout)
|
| 77 |
-
layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead,
|
| 78 |
-
dim_feedforward=2048, batch_first=True,
|
| 79 |
-
dropout=dropout, activation='gelu')
|
| 80 |
-
self.dec = nn.TransformerDecoder(layer, num_layers=num_layers)
|
| 81 |
-
self.fc = nn.Linear(d_model, vocab_size)
|
| 82 |
-
|
| 83 |
-
@staticmethod
|
| 84 |
-
def _causal_mask(L, device, dtype):
|
| 85 |
-
m = torch.full((L, L), float('-inf'), device=device, dtype=dtype)
|
| 86 |
-
return torch.triu(m, diagonal=1)
|
| 87 |
-
|
| 88 |
-
def forward(self, tgt_tokens: torch.Tensor, memory: torch.Tensor) -> torch.Tensor:
|
| 89 |
-
"""
|
| 90 |
-
tgt_tokens: (B, T), memory: (B, S, D)
|
| 91 |
-
returns logits: (B, T, V)
|
| 92 |
-
"""
|
| 93 |
-
B, T = tgt_tokens.shape
|
| 94 |
-
x = self.embed(tgt_tokens) # (B,T,D)
|
| 95 |
-
x = self.pos(x)
|
| 96 |
-
tgt_mask = self._causal_mask(T, x.device, x.dtype) # (T,T) float mask
|
| 97 |
-
tgt_kpm = (tgt_tokens == self.pad_idx) # (B,T) bool
|
| 98 |
-
y = self.dec(tgt=x, memory=memory,
|
| 99 |
-
tgt_mask=tgt_mask,
|
| 100 |
-
tgt_key_padding_mask=tgt_kpm)
|
| 101 |
-
return self.fc(y) # (B,T,V)
|
| 102 |
-
|
| 103 |
-
# ---------- full model ----------
|
| 104 |
-
class Net(nn.Module):
|
| 105 |
-
"""
|
| 106 |
-
Returns a Tensor from forward() (never a tuple), so metrics like BLEU can call .dim().
|
| 107 |
-
API:
|
| 108 |
-
- __init__(in_shape, out_shape, prm, device)
|
| 109 |
-
- forward(images, captions=None) -> logits Tensor
|
| 110 |
-
- train_setup(prm)
|
| 111 |
-
- learn(train_data)
|
| 112 |
-
"""
|
| 113 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 114 |
-
super().__init__()
|
| 115 |
-
self.device = device
|
| 116 |
-
self.in_channels = int(in_shape[1])
|
| 117 |
-
self.vocab_size = int(out_shape[0])
|
| 118 |
-
|
| 119 |
-
# Hyperparams
|
| 120 |
-
self.dropout_p = float(prm.get('dropout', 0.1))
|
| 121 |
-
self.max_len = int(prm.get('max_len', 20))
|
| 122 |
-
|
| 123 |
-
# Special tokens
|
| 124 |
-
self.pad_idx, self.bos_id, self.eos_id = _special_ids(GLOBAL_CAPTION_VOCAB or {}, self.vocab_size)
|
| 125 |
-
|
| 126 |
-
# Encoder -> sequence of d_model features
|
| 127 |
-
d_model = 512
|
| 128 |
-
enc_feat_ch = 256 # AlexNet conv5 output channels
|
| 129 |
-
self.encoder = CNNEncoder(self.in_channels, feat_ch=enc_feat_ch)
|
| 130 |
-
self.enc_proj = nn.Linear(enc_feat_ch, d_model)
|
| 131 |
-
self.enc_pos = PositionalEncoding(d_model, dropout=self.dropout_p)
|
| 132 |
-
self.enc_drop = nn.Dropout(self.dropout_p)
|
| 133 |
-
|
| 134 |
-
# Transformer decoder
|
| 135 |
-
self.decoder = TransformerDecoder(self.vocab_size, d_model=d_model, nhead=8,
|
| 136 |
-
num_layers=4, dropout=self.dropout_p, pad_idx=self.pad_idx)
|
| 137 |
-
|
| 138 |
-
# Training attrs init in train_setup
|
| 139 |
-
self.criteria = None
|
| 140 |
-
self.optimizer = None
|
| 141 |
-
self.scaler = None
|
| 142 |
-
|
| 143 |
-
# -- encoder helper --
|
| 144 |
-
def _encode(self, images: torch.Tensor) -> torch.Tensor:
|
| 145 |
-
f = self.encoder(images) # (B,C,h,w) with C=256
|
| 146 |
-
B, C, h, w = f.shape
|
| 147 |
-
seq = f.view(B, C, h*w).permute(0, 2, 1) # (B,S,C) where S=h*w
|
| 148 |
-
seq = self.enc_proj(seq) # (B,S,D)
|
| 149 |
-
seq = self.enc_pos(seq)
|
| 150 |
-
seq = self.enc_drop(seq)
|
| 151 |
-
return seq # (B,S,D)
|
| 152 |
-
|
| 153 |
-
# -- forward --
|
| 154 |
-
def forward(self, images, captions=None):
|
| 155 |
-
"""
|
| 156 |
-
Training (teacher forcing):
|
| 157 |
-
inputs = captions[:, :-1] -> logits over positions 1..T-1
|
| 158 |
-
returns logits: (B, T-1, V)
|
| 159 |
-
Inference (captions=None):
|
| 160 |
-
greedy decode up to max_len
|
| 161 |
-
returns logits: (B, L, V) of generated steps
|
| 162 |
-
"""
|
| 163 |
-
assert images.dim() == 4, "images must be (B,C,H,W)"
|
| 164 |
-
memory = self._encode(images) # (B,S,D)
|
| 165 |
-
|
| 166 |
-
if captions is not None:
|
| 167 |
-
if captions.ndim == 3:
|
| 168 |
-
captions = captions[:, 0, :] # (B,T)
|
| 169 |
-
inputs = captions[:, :-1] # (B,T-1)
|
| 170 |
-
logits = self.decoder(inputs, memory) # (B,T-1,V)
|
| 171 |
-
return logits # Tensor ONLY
|
| 172 |
-
|
| 173 |
-
# Inference: greedy
|
| 174 |
-
B = images.size(0)
|
| 175 |
-
device = images.device
|
| 176 |
-
cur = torch.full((B, 1), self.bos_id, dtype=torch.long, device=device)
|
| 177 |
-
steps = []
|
| 178 |
-
for _ in range(self.max_len):
|
| 179 |
-
step_logits = self.decoder(cur, memory)[:, -1:, :] # (B,1,V)
|
| 180 |
-
steps.append(step_logits)
|
| 181 |
-
next_tok = step_logits.argmax(dim=-1) # (B,1)
|
| 182 |
-
cur = torch.cat([cur, next_tok], dim=1)
|
| 183 |
-
if (next_tok.squeeze(1) == self.eos_id).all():
|
| 184 |
-
break
|
| 185 |
-
logits = torch.cat(steps, dim=1) if steps else torch.zeros((B, 0, self.vocab_size), device=device)
|
| 186 |
-
return logits # Tensor ONLY
|
| 187 |
-
|
| 188 |
-
# -- training setup --
|
| 189 |
-
def train_setup(self, prm):
|
| 190 |
-
self.to(self.device)
|
| 191 |
-
# Loss: ignore PAD; label smoothing helps BLEU
|
| 192 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=self.pad_idx, label_smoothing=0.1).to(self.device),)
|
| 193 |
-
# Map "momentum" → AdamW beta1 to keep CLI semantics
|
| 194 |
-
beta1 = float(prm.get('momentum', 0.9))
|
| 195 |
-
self.optimizer = torch.optim.AdamW(self.parameters(),
|
| 196 |
-
lr=float(prm['lr']),
|
| 197 |
-
betas=(beta1, 0.999),
|
| 198 |
-
weight_decay=1e-4)
|
| 199 |
-
# AMP (PyTorch 2.x API)
|
| 200 |
-
self.scaler = torch.amp.GradScaler('cuda', enabled=(self.device.type == 'cuda'))
|
| 201 |
-
|
| 202 |
-
# -- one epoch training loop --
|
| 203 |
-
def learn(self, train_data):
|
| 204 |
-
"""
|
| 205 |
-
Expects batches like (images, captions, *rest).
|
| 206 |
-
"""
|
| 207 |
-
assert self.criteria and self.optimizer is not None and self.scaler is not None, "Call train_setup(prm) first."
|
| 208 |
-
self.train()
|
| 209 |
-
amp_device = 'cuda' if self.device.type == 'cuda' else 'cpu'
|
| 210 |
-
for batch in train_data:
|
| 211 |
-
if isinstance(batch, (list, tuple)):
|
| 212 |
-
images, captions = batch[0], batch[1]
|
| 213 |
-
else:
|
| 214 |
-
images, captions = batch
|
| 215 |
-
images = images.to(self.device, non_blocking=True)
|
| 216 |
-
captions = captions.to(self.device, non_blocking=True)
|
| 217 |
-
|
| 218 |
-
with torch.amp.autocast(amp_device, enabled=(self.device.type == 'cuda')):
|
| 219 |
-
if captions.ndim == 3:
|
| 220 |
-
captions = captions[:, 0, :]
|
| 221 |
-
logits = self.forward(images, captions) # (B,T-1,V)
|
| 222 |
-
targets = captions[:, 1:] # (B,T-1)
|
| 223 |
-
loss = self.criteria[0](logits.reshape(-1, logits.size(-1)),
|
| 224 |
-
targets.reshape(-1))
|
| 225 |
-
|
| 226 |
-
self.optimizer.zero_grad(set_to_none=True)
|
| 227 |
-
self.scaler.scale(loss).backward()
|
| 228 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 229 |
-
self.scaler.step(self.optimizer)
|
| 230 |
-
self.scaler.update()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C10C-ResNetTransformer-187ccbee8050ac295637ecedecb4da1e.py
DELETED
|
@@ -1,193 +0,0 @@
|
|
| 1 |
-
import math
|
| 2 |
-
from typing import Optional
|
| 3 |
-
|
| 4 |
-
import torch
|
| 5 |
-
import torch.nn as nn
|
| 6 |
-
import torch.nn.functional as F
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
def supported_hyperparameters():
|
| 10 |
-
return {"lr", "momentum"}
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
# ---------------- blocks ----------------
|
| 14 |
-
|
| 15 |
-
class SEBlock(nn.Module):
|
| 16 |
-
def __init__(self, c: int, r: int = 8):
|
| 17 |
-
super().__init__()
|
| 18 |
-
m = max(4, c // r)
|
| 19 |
-
self.fc1 = nn.Linear(c, m)
|
| 20 |
-
self.fc2 = nn.Linear(m, c)
|
| 21 |
-
|
| 22 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 23 |
-
b, c, h, w = x.size()
|
| 24 |
-
s = x.mean(dim=(2, 3))
|
| 25 |
-
s = F.relu(self.fc1(s))
|
| 26 |
-
s = torch.sigmoid(self.fc2(s)).view(b, c, 1, 1)
|
| 27 |
-
return x * s
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
class ConvBlock(nn.Module):
|
| 31 |
-
def __init__(self, in_c: int, out_c: int, stride: int = 1):
|
| 32 |
-
super().__init__()
|
| 33 |
-
self.conv1 = nn.Conv2d(in_c, out_c, 3, stride=stride, padding=1, bias=False)
|
| 34 |
-
self.bn1 = nn.BatchNorm2d(out_c)
|
| 35 |
-
self.conv2 = nn.Conv2d(out_c, out_c, 3, padding=1, bias=False)
|
| 36 |
-
self.bn2 = nn.BatchNorm2d(out_c)
|
| 37 |
-
self.se = SEBlock(out_c)
|
| 38 |
-
self.skip = None
|
| 39 |
-
if stride != 1 or in_c != out_c:
|
| 40 |
-
self.skip = nn.Sequential(nn.Conv2d(in_c, out_c, 1, stride=stride, bias=False),
|
| 41 |
-
nn.BatchNorm2d(out_c))
|
| 42 |
-
|
| 43 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 44 |
-
id = x
|
| 45 |
-
x = F.relu(self.bn1(self.conv1(x)), inplace=True)
|
| 46 |
-
x = self.bn2(self.conv2(x))
|
| 47 |
-
if self.skip is not None:
|
| 48 |
-
id = self.skip(id)
|
| 49 |
-
x = F.relu(x + id, inplace=True)
|
| 50 |
-
x = self.se(x)
|
| 51 |
-
return x
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
class PositionalEncodingBF(nn.Module):
|
| 55 |
-
def __init__(self, d_model: int, max_len: int = 512):
|
| 56 |
-
super().__init__()
|
| 57 |
-
pe = torch.zeros(max_len, d_model)
|
| 58 |
-
pos = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
|
| 59 |
-
div = torch.exp(torch.arange(0, d_model, 2, dtype=torch.float) * (-math.log(10000.0) / d_model))
|
| 60 |
-
pe[:, 0::2] = torch.sin(pos * div)
|
| 61 |
-
pe[:, 1::2] = torch.cos(pos * div)
|
| 62 |
-
self.register_buffer("pe", pe, persistent=False)
|
| 63 |
-
|
| 64 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 65 |
-
# x: [B, T, D]
|
| 66 |
-
T = x.size(1)
|
| 67 |
-
return x + self.pe[:T].unsqueeze(0)
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
# ---------------- encoder/decoder ----------------
|
| 71 |
-
|
| 72 |
-
class CNNEncoder(nn.Module):
|
| 73 |
-
def __init__(self, in_ch: int, d_model: int):
|
| 74 |
-
super().__init__()
|
| 75 |
-
self.stem = nn.Sequential(
|
| 76 |
-
nn.Conv2d(in_ch, 64, 7, stride=2, padding=3, bias=False),
|
| 77 |
-
nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(3, stride=2, padding=1),
|
| 78 |
-
)
|
| 79 |
-
self.s1 = ConvBlock(64, 128, stride=2)
|
| 80 |
-
self.s2 = ConvBlock(128, 256, stride=2)
|
| 81 |
-
self.s3 = ConvBlock(256, 256, stride=1)
|
| 82 |
-
self.head = nn.Sequential(
|
| 83 |
-
nn.Conv2d(256, d_model, 1, bias=False), nn.BatchNorm2d(d_model), nn.ReLU(inplace=True),
|
| 84 |
-
nn.AdaptiveAvgPool2d(1),
|
| 85 |
-
)
|
| 86 |
-
|
| 87 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 88 |
-
x = self.stem(x); x = self.s1(x); x = self.s2(x); x = self.s3(x)
|
| 89 |
-
x = self.head(x).squeeze(-1).squeeze(-1) # [B, D]
|
| 90 |
-
return x.unsqueeze(1) # [B, 1, D]
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
class TransformerCaptionDecoder(nn.Module):
|
| 94 |
-
def __init__(self, vocab: int, d_model: int = 640, nhead: int = 8, layers: int = 2, dim_ff: int = 2048, dropout: float = 0.2):
|
| 95 |
-
super().__init__()
|
| 96 |
-
assert d_model % nhead == 0
|
| 97 |
-
self.embed = nn.Embedding(vocab, d_model, padding_idx=0)
|
| 98 |
-
self.pe = PositionalEncodingBF(d_model)
|
| 99 |
-
layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead, dim_feedforward=dim_ff,
|
| 100 |
-
dropout=dropout, batch_first=True)
|
| 101 |
-
self.dec = nn.TransformerDecoder(layer, num_layers=layers)
|
| 102 |
-
self.proj = nn.Linear(d_model, vocab, bias=False)
|
| 103 |
-
self.proj.weight = self.embed.weight
|
| 104 |
-
|
| 105 |
-
@staticmethod
|
| 106 |
-
def _causal_mask(T: int, device: torch.device):
|
| 107 |
-
m = torch.full((T, T), float("-inf"), device=device)
|
| 108 |
-
return torch.triu(m, diagonal=1)
|
| 109 |
-
|
| 110 |
-
def forward(self, tokens: torch.Tensor, memory: torch.Tensor) -> torch.Tensor:
|
| 111 |
-
x = self.embed(tokens) * math.sqrt(self.embed.embedding_dim)
|
| 112 |
-
x = self.pe(x)
|
| 113 |
-
mask = self._causal_mask(x.size(1), x.device)
|
| 114 |
-
x = self.dec(tgt=x, memory=memory, tgt_mask=mask)
|
| 115 |
-
return self.proj(x) # [B, T, V]
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
# ---------------- Net (API) ----------------
|
| 119 |
-
|
| 120 |
-
class Net(nn.Module):
|
| 121 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device):
|
| 122 |
-
super().__init__()
|
| 123 |
-
self.device = device
|
| 124 |
-
in_ch = int(in_shape[1])
|
| 125 |
-
vocab = int(out_shape[0])
|
| 126 |
-
|
| 127 |
-
d_model = int(prm.get("hidden_dim", 640))
|
| 128 |
-
nhead = int(prm.get("nhead", 8))
|
| 129 |
-
layers = int(prm.get("dec_layers", 2))
|
| 130 |
-
dim_ff = int(prm.get("dim_ff", 2048))
|
| 131 |
-
dropout = float(prm.get("dropout", 0.2))
|
| 132 |
-
|
| 133 |
-
self.encoder = CNNEncoder(in_ch, d_model)
|
| 134 |
-
self.rnn = TransformerCaptionDecoder(vocab, d_model, nhead, layers, dim_ff, dropout)
|
| 135 |
-
self.vocab = vocab
|
| 136 |
-
|
| 137 |
-
self.criterion = nn.CrossEntropyLoss(ignore_index=0, label_smoothing=0.05)
|
| 138 |
-
self.optimizer = None
|
| 139 |
-
|
| 140 |
-
@staticmethod
|
| 141 |
-
def _norm_caps(caps: Optional[torch.Tensor]) -> Optional[torch.Tensor]:
|
| 142 |
-
if caps is None: return None
|
| 143 |
-
if caps.ndim == 1: caps = caps.unsqueeze(0)
|
| 144 |
-
elif caps.ndim == 3: caps = caps[:, 0, :]
|
| 145 |
-
return caps.long()
|
| 146 |
-
|
| 147 |
-
def train_setup(self, prm: dict):
|
| 148 |
-
self.to(self.device)
|
| 149 |
-
lr = max(float(prm.get("lr", 1e-3)), 1e-3)
|
| 150 |
-
b1 = min(0.99, max(0.7, float(prm.get("momentum", 0.9))))
|
| 151 |
-
self.optimizer = torch.optim.AdamW(self.parameters(), lr=lr, betas=(b1, 0.999), weight_decay=1e-4)
|
| 152 |
-
self.criterion = self.criterion.to(self.device)
|
| 153 |
-
|
| 154 |
-
def learn(self, train_data):
|
| 155 |
-
self.train()
|
| 156 |
-
for images, captions in train_data:
|
| 157 |
-
images = images.to(self.device, non_blocking=True)
|
| 158 |
-
captions = captions.to(self.device, non_blocking=True)
|
| 159 |
-
|
| 160 |
-
caps = self._norm_caps(captions) # [B, T]
|
| 161 |
-
inp, tgt = caps[:, :-1], caps[:, 1:] # [B, T-1]
|
| 162 |
-
|
| 163 |
-
mem = self.encoder(images) # [B, 1, D]
|
| 164 |
-
logits = self.rnn(inp, mem) # [B, T-1, V]
|
| 165 |
-
|
| 166 |
-
assert logits.shape[1] == inp.shape[1] and logits.shape[-1] == self.vocab
|
| 167 |
-
loss = self.criterion(logits.reshape(-1, self.vocab), tgt.reshape(-1))
|
| 168 |
-
|
| 169 |
-
self.optimizer.zero_grad(set_to_none=True)
|
| 170 |
-
loss.backward()
|
| 171 |
-
torch.nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 172 |
-
self.optimizer.step()
|
| 173 |
-
|
| 174 |
-
def forward(self, images: torch.Tensor, captions: Optional[torch.Tensor] = None, hidden_state=None) -> torch.Tensor:
|
| 175 |
-
images = images.to(self.device, non_blocking=True)
|
| 176 |
-
mem = self.encoder(images) # [B, 1, D]
|
| 177 |
-
|
| 178 |
-
if captions is None:
|
| 179 |
-
# simple greedy stub
|
| 180 |
-
B = images.size(0)
|
| 181 |
-
seq = torch.full((B, 1), 1, dtype=torch.long, device=self.device) # <SOS>=1
|
| 182 |
-
for _ in range(19):
|
| 183 |
-
lg = self.rnn(seq, mem)
|
| 184 |
-
nxt = lg[:, -1, :].argmax(-1, keepdim=True)
|
| 185 |
-
seq = torch.cat([seq, nxt], dim=1)
|
| 186 |
-
if (nxt == 2).all(): break # <EOS>=2
|
| 187 |
-
return self.rnn(seq, mem)
|
| 188 |
-
|
| 189 |
-
caps = self._norm_caps(captions).to(self.device) # [B, T]
|
| 190 |
-
inp = caps[:, :-1] # [B, T-1]
|
| 191 |
-
logits = self.rnn(inp, mem) # [B, T-1, V]
|
| 192 |
-
pad = torch.zeros((caps.size(0), 1, self.vocab), device=logits.device, dtype=logits.dtype)
|
| 193 |
-
return torch.cat([pad, logits], dim=1) # [B, T, V]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C5C-RESNETLSTM-4.py
DELETED
|
@@ -1,222 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
from torch import nn, Tensor
|
| 3 |
-
from typing import Any, Optional
|
| 4 |
-
from collections import Counter
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
# NN-GPT / NN-Dataset expect exactly {'lr','momentum'} at module level
|
| 9 |
-
return {"lr", "momentum"}
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
def _first_int(x: Any) -> int:
|
| 13 |
-
if isinstance(x, int):
|
| 14 |
-
return x
|
| 15 |
-
if isinstance(x, (tuple, list)) and len(x) > 0:
|
| 16 |
-
return _first_int(x[0])
|
| 17 |
-
try:
|
| 18 |
-
return int(x)
|
| 19 |
-
except Exception:
|
| 20 |
-
return 10000
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
class Net(nn.Module):
|
| 24 |
-
def __init__(self, in_shape: Any, out_shape: Any, prm: dict, device: torch.device, *_, **__):
|
| 25 |
-
super().__init__()
|
| 26 |
-
|
| 27 |
-
self.device = device
|
| 28 |
-
self.in_shape = in_shape
|
| 29 |
-
self.out_shape = out_shape
|
| 30 |
-
self.prm = dict(prm) if prm is not None else {}
|
| 31 |
-
|
| 32 |
-
# Infer channels from in_shape (supports (C,H,W) or (N,C,H,W))
|
| 33 |
-
if isinstance(in_shape, (tuple, list)) and len(in_shape) > 1:
|
| 34 |
-
self.in_channels = int(in_shape[1])
|
| 35 |
-
else:
|
| 36 |
-
self.in_channels = 3
|
| 37 |
-
|
| 38 |
-
# vocab_size from out_shape, e.g. (V,) or V
|
| 39 |
-
self.vocab_size = _first_int(out_shape)
|
| 40 |
-
|
| 41 |
-
emb_dim = 512
|
| 42 |
-
hid_dim = 512
|
| 43 |
-
drop = float(self.prm.get("dropout", 0.2))
|
| 44 |
-
|
| 45 |
-
# Stable CNN encoder -> [B, 256]
|
| 46 |
-
self.encoder = nn.Sequential(
|
| 47 |
-
nn.Conv2d(self.in_channels, 64, 3, 2, 1),
|
| 48 |
-
nn.ReLU(inplace=True),
|
| 49 |
-
nn.Conv2d(64, 128, 3, 2, 1),
|
| 50 |
-
nn.ReLU(inplace=True),
|
| 51 |
-
nn.Conv2d(128, 256, 3, 2, 1),
|
| 52 |
-
nn.ReLU(inplace=True),
|
| 53 |
-
nn.AdaptiveAvgPool2d((1, 1)),
|
| 54 |
-
nn.Flatten()
|
| 55 |
-
)
|
| 56 |
-
self.enc_fc = nn.Linear(256, emb_dim)
|
| 57 |
-
|
| 58 |
-
# Caption decoder
|
| 59 |
-
self.embed = nn.Embedding(self.vocab_size, emb_dim, padding_idx=0)
|
| 60 |
-
self.drop = nn.Dropout(drop)
|
| 61 |
-
self.lstm = nn.LSTM(emb_dim, hid_dim, batch_first=True)
|
| 62 |
-
self.fc = nn.Linear(hid_dim, self.vocab_size)
|
| 63 |
-
|
| 64 |
-
# Training helpers
|
| 65 |
-
self.criterion: Optional[nn.Module] = None
|
| 66 |
-
self.optimizer: Optional[torch.optim.Optimizer] = None
|
| 67 |
-
|
| 68 |
-
# Token stats for a simple fallback in predict()
|
| 69 |
-
self._token_counts = Counter()
|
| 70 |
-
self._have_stats = False
|
| 71 |
-
self._bos = 1
|
| 72 |
-
self._eos = 2
|
| 73 |
-
self._pad = 0
|
| 74 |
-
self._max_len = 16
|
| 75 |
-
|
| 76 |
-
# Class-level helper (not used by harness, but kept)
|
| 77 |
-
def supported_hyperparameters(self):
|
| 78 |
-
return {"lr", "momentum", "dropout"}
|
| 79 |
-
|
| 80 |
-
def _norm(self, caps: Tensor) -> Tensor:
|
| 81 |
-
# Normalize caption shape to [B, T]
|
| 82 |
-
if caps.dim() == 1:
|
| 83 |
-
return caps.unsqueeze(0)
|
| 84 |
-
if caps.dim() == 3:
|
| 85 |
-
# e.g. [B, 1, T]
|
| 86 |
-
return caps[:, 0, :]
|
| 87 |
-
return caps
|
| 88 |
-
|
| 89 |
-
def _enc(self, x: Tensor):
|
| 90 |
-
# Encode image -> initial LSTM hidden state
|
| 91 |
-
feats = self.encoder(x) # [B, 256]
|
| 92 |
-
ctx = self.enc_fc(feats) # [B, emb_dim]
|
| 93 |
-
h0 = torch.tanh(ctx).unsqueeze(0) # [1, B, H]
|
| 94 |
-
c0 = torch.tanh(ctx).unsqueeze(0) # [1, B, H]
|
| 95 |
-
return (h0, c0)
|
| 96 |
-
|
| 97 |
-
def forward(self, images: Tensor, captions: Optional[Tensor] = None):
|
| 98 |
-
images = images.to(self.device, dtype=torch.float32)
|
| 99 |
-
|
| 100 |
-
# Training / teacher forcing path
|
| 101 |
-
if captions is not None:
|
| 102 |
-
captions = captions.to(self.device, dtype=torch.long)
|
| 103 |
-
captions = self._norm(captions) # [B, T]
|
| 104 |
-
|
| 105 |
-
if captions.size(1) <= 1:
|
| 106 |
-
# Degenerate case: no real caption content
|
| 107 |
-
B = captions.size(0)
|
| 108 |
-
dummy = torch.zeros(B, 1, self.lstm.hidden_size, device=self.device)
|
| 109 |
-
return self.fc(dummy)
|
| 110 |
-
|
| 111 |
-
# Update frequency stats for predict() fallback
|
| 112 |
-
with torch.no_grad():
|
| 113 |
-
valid = captions[captions != self._pad].reshape(-1)
|
| 114 |
-
for t in valid.tolist():
|
| 115 |
-
self._token_counts[int(t)] += 1
|
| 116 |
-
self._have_stats = len(self._token_counts) > 0
|
| 117 |
-
|
| 118 |
-
dec_in = captions[:, :-1] # [B, T-1]
|
| 119 |
-
emb = self.drop(self.embed(dec_in)) # [B, T-1, E]
|
| 120 |
-
h0, c0 = self._enc(images) # ([1,B,H],[1,B,H])
|
| 121 |
-
out, _ = self.lstm(emb, (h0, c0)) # [B, T-1, H]
|
| 122 |
-
logits = self.fc(self.drop(out)) # [B, T-1, V]
|
| 123 |
-
return logits
|
| 124 |
-
|
| 125 |
-
# Inference path: generate tokens for BLEU
|
| 126 |
-
return self.predict(images)
|
| 127 |
-
|
| 128 |
-
def train_setup(self, prm: dict):
|
| 129 |
-
lr = float(prm.get("lr", 1e-3))
|
| 130 |
-
mom = float(prm.get("momentum", 0.9))
|
| 131 |
-
drop = float(prm.get("dropout", self.prm.get("dropout", 0.2)))
|
| 132 |
-
self.drop.p = drop
|
| 133 |
-
|
| 134 |
-
self.to(self.device)
|
| 135 |
-
self.train()
|
| 136 |
-
|
| 137 |
-
self.criterion = nn.CrossEntropyLoss(ignore_index=self._pad)
|
| 138 |
-
self.optimizer = torch.optim.AdamW(self.parameters(), lr=lr, betas=(mom, 0.999))
|
| 139 |
-
|
| 140 |
-
def learn(self, data):
|
| 141 |
-
if self.optimizer is None:
|
| 142 |
-
prm = getattr(data, "prm", self.prm)
|
| 143 |
-
self.train_setup(prm)
|
| 144 |
-
|
| 145 |
-
self.train()
|
| 146 |
-
|
| 147 |
-
for batch in data:
|
| 148 |
-
if isinstance(batch, (list, tuple)):
|
| 149 |
-
if len(batch) < 2:
|
| 150 |
-
continue
|
| 151 |
-
imgs, caps = batch[0], batch[1]
|
| 152 |
-
elif isinstance(batch, dict):
|
| 153 |
-
imgs = batch.get("x", None)
|
| 154 |
-
caps = batch.get("y", None)
|
| 155 |
-
if imgs is None or caps is None:
|
| 156 |
-
continue
|
| 157 |
-
else:
|
| 158 |
-
imgs = getattr(batch, "x", None)
|
| 159 |
-
caps = getattr(batch, "y", None)
|
| 160 |
-
if imgs is None or caps is None:
|
| 161 |
-
continue
|
| 162 |
-
|
| 163 |
-
imgs = imgs.to(self.device)
|
| 164 |
-
caps = caps.to(self.device)
|
| 165 |
-
caps = self._norm(caps)
|
| 166 |
-
if caps.size(1) <= 1:
|
| 167 |
-
continue
|
| 168 |
-
|
| 169 |
-
logits = self.forward(imgs, caps) # [B, T-1, V]
|
| 170 |
-
targets = caps[:, 1:] # [B, T-1]
|
| 171 |
-
|
| 172 |
-
loss = self.criterion(
|
| 173 |
-
logits.reshape(-1, self.vocab_size),
|
| 174 |
-
targets.reshape(-1),
|
| 175 |
-
)
|
| 176 |
-
|
| 177 |
-
self.optimizer.zero_grad(set_to_none=True)
|
| 178 |
-
loss.backward()
|
| 179 |
-
nn.utils.clip_grad_norm_(self.parameters(), 1.0)
|
| 180 |
-
self.optimizer.step()
|
| 181 |
-
|
| 182 |
-
@torch.no_grad()
|
| 183 |
-
def predict(self, images: Tensor) -> Tensor:
|
| 184 |
-
"""
|
| 185 |
-
Greedy decoding for BLEU eval.
|
| 186 |
-
Returns [B, T] token IDs.
|
| 187 |
-
"""
|
| 188 |
-
self.eval()
|
| 189 |
-
images = images.to(self.device)
|
| 190 |
-
B = images.size(0)
|
| 191 |
-
|
| 192 |
-
# If we have token stats from training, return a simple "common tokens" caption
|
| 193 |
-
if self._have_stats:
|
| 194 |
-
common = [
|
| 195 |
-
t for (t, _) in self._token_counts.most_common(self._max_len + 4)
|
| 196 |
-
if t != self._pad
|
| 197 |
-
]
|
| 198 |
-
if not common:
|
| 199 |
-
common = [self._bos]
|
| 200 |
-
|
| 201 |
-
base = common[: self._max_len - 2]
|
| 202 |
-
seq = [self._bos] + base + [self._eos]
|
| 203 |
-
tokens = torch.tensor(seq, dtype=torch.long, device=self.device)
|
| 204 |
-
return tokens.unsqueeze(0).repeat(B, 1)
|
| 205 |
-
|
| 206 |
-
# Otherwise, decode with LSTM
|
| 207 |
-
h0, c0 = self._enc(images)
|
| 208 |
-
tokens = torch.full((B, 1), self._bos, dtype=torch.long, device=self.device)
|
| 209 |
-
|
| 210 |
-
for _ in range(self._max_len - 1):
|
| 211 |
-
emb = self.drop(self.embed(tokens[:, -1:])) # [B,1,E]
|
| 212 |
-
out, (h0, c0) = self.lstm(emb, (h0, c0)) # [B,1,H]
|
| 213 |
-
nxt = self.fc(out).argmax(-1) # [B,1]
|
| 214 |
-
tokens = torch.cat([tokens, nxt], dim=1)
|
| 215 |
-
if (nxt == self._eos).all():
|
| 216 |
-
break
|
| 217 |
-
|
| 218 |
-
return tokens
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
def model_net(in_shape, out_shape, prm, device):
|
| 222 |
-
return Net(in_shape, out_shape, prm, device)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C5C-RESNETLSTM-c42512d71480c8ef10f31e3e6c33bbdf.py
DELETED
|
@@ -1,150 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
|
| 4 |
-
def supported_hyperparameters():
|
| 5 |
-
return {'lr', 'momentum'}
|
| 6 |
-
|
| 7 |
-
class Encoder(nn.Module):
|
| 8 |
-
def __init__(self, in_channels: int, embed_size: int):
|
| 9 |
-
super().__init__()
|
| 10 |
-
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=7, stride=2, padding=3)
|
| 11 |
-
self.bn1 = nn.BatchNorm2d(32)
|
| 12 |
-
self.relu = nn.ReLU(inplace=True)
|
| 13 |
-
self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 14 |
-
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
|
| 15 |
-
self.bn2 = nn.BatchNorm2d(64)
|
| 16 |
-
self.conv3 = nn.Conv2d(64, 128, kernel_size=3, padding=1)
|
| 17 |
-
self.bn3 = nn.BatchNorm2d(128)
|
| 18 |
-
self.adap_pool = nn.AdaptiveMaxPool2d(7)
|
| 19 |
-
self.fc = nn.Linear(128 * 7 * 7, embed_size)
|
| 20 |
-
|
| 21 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 22 |
-
x = self.relu(self.bn1(self.conv1(x)))
|
| 23 |
-
x = self.pool(x)
|
| 24 |
-
x = self.relu(self.bn2(self.conv2(x)))
|
| 25 |
-
x = self.pool(self.relu(self.bn3(self.conv3(x))))
|
| 26 |
-
x = self.adap_pool(x)
|
| 27 |
-
x = x.view(x.size(0), -1)
|
| 28 |
-
x = self.fc(x)
|
| 29 |
-
return x
|
| 30 |
-
|
| 31 |
-
class DecoderRNN(nn.Module):
|
| 32 |
-
def __init__(self, embed_size: int, hidden_size: int, vocab_size: int, drop_prob: float = 0.0, num_layers: int = 1):
|
| 33 |
-
super().__init__()
|
| 34 |
-
self.embed = nn.Embedding(vocab_size, embed_size)
|
| 35 |
-
self.lstm = nn.LSTM(input_size=embed_size * 2, hidden_size=hidden_size, num_layers=num_layers, batch_first=True, dropout=0.0)
|
| 36 |
-
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 37 |
-
self.hidden_size = hidden_size
|
| 38 |
-
self.embed_size = embed_size
|
| 39 |
-
self.vocab_size = vocab_size
|
| 40 |
-
|
| 41 |
-
def init_hidden(self, batch_size: int, device: torch.device):
|
| 42 |
-
h = torch.zeros(1, batch_size, self.hidden_size, device=device)
|
| 43 |
-
c = torch.zeros(1, batch_size, self.hidden_size, device=device)
|
| 44 |
-
return (h, c)
|
| 45 |
-
|
| 46 |
-
def init_zero_hidden(self, batch_size: int, device: torch.device):
|
| 47 |
-
return self.init_hidden(batch_size, device)
|
| 48 |
-
|
| 49 |
-
def forward(self, inputs: torch.Tensor, hidden: tuple, encoder_features: torch.Tensor):
|
| 50 |
-
B, T = inputs.size()
|
| 51 |
-
embeds = self.embed(inputs)
|
| 52 |
-
h, c = hidden
|
| 53 |
-
outputs = []
|
| 54 |
-
for t in range(T):
|
| 55 |
-
word_t = embeds[:, t, :]
|
| 56 |
-
step_in = torch.cat([word_t, encoder_features], dim=1)
|
| 57 |
-
out, (h, c) = self.lstm(step_in.unsqueeze(1), (h, c))
|
| 58 |
-
outputs.append(self.fc(out.squeeze(1)))
|
| 59 |
-
logits = torch.stack(outputs, dim=1)
|
| 60 |
-
return logits, (h, c)
|
| 61 |
-
|
| 62 |
-
@torch.no_grad()
|
| 63 |
-
def generate(self, encoder_features: torch.Tensor, hidden: tuple, max_len: int = 50, start_token: int = 1, end_token: int = 2):
|
| 64 |
-
device = encoder_features.device
|
| 65 |
-
B = encoder_features.size(0)
|
| 66 |
-
cur = torch.full((B,), start_token, dtype=torch.long, device=device)
|
| 67 |
-
generated = []
|
| 68 |
-
for _ in range(max_len):
|
| 69 |
-
step_logits, hidden = self.forward(cur.unsqueeze(1), hidden, encoder_features)
|
| 70 |
-
next_ids = step_logits.squeeze(1).argmax(dim=1)
|
| 71 |
-
generated.append(next_ids)
|
| 72 |
-
cur = next_ids
|
| 73 |
-
if (next_ids == end_token).all():
|
| 74 |
-
break
|
| 75 |
-
if len(generated) == 0:
|
| 76 |
-
return torch.empty(B, 0, dtype=torch.long, device=device), hidden
|
| 77 |
-
return torch.stack(generated, dim=1), hidden
|
| 78 |
-
|
| 79 |
-
class Net(nn.Module):
|
| 80 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 81 |
-
super().__init__()
|
| 82 |
-
self.device = device
|
| 83 |
-
in_channels = int(in_shape[1])
|
| 84 |
-
self.vocab_size = int(out_shape[0])
|
| 85 |
-
embed_size = 256
|
| 86 |
-
hidden_size = 512
|
| 87 |
-
self.encoder = Encoder(in_channels, embed_size)
|
| 88 |
-
self.rnn = DecoderRNN(embed_size, hidden_size, self.vocab_size, drop_prob=0.0)
|
| 89 |
-
self.criterion = nn.CrossEntropyLoss(ignore_index=0)
|
| 90 |
-
self.optimizer = None
|
| 91 |
-
|
| 92 |
-
@staticmethod
|
| 93 |
-
def _normalize_captions(captions: torch.Tensor) -> torch.Tensor:
|
| 94 |
-
if captions is None:
|
| 95 |
-
return None
|
| 96 |
-
if captions.ndim == 3:
|
| 97 |
-
captions = captions[:, 0, :]
|
| 98 |
-
elif captions.ndim == 1:
|
| 99 |
-
captions = captions.unsqueeze(0)
|
| 100 |
-
if captions.dtype != torch.long:
|
| 101 |
-
captions = captions.long()
|
| 102 |
-
return captions
|
| 103 |
-
|
| 104 |
-
def _ensure_hidden(self, hidden_state, batch_size: int):
|
| 105 |
-
if hidden_state is None:
|
| 106 |
-
return self.rnn.init_zero_hidden(batch_size, self.device)
|
| 107 |
-
h = hidden_state[0]
|
| 108 |
-
if h.size(1) != batch_size:
|
| 109 |
-
return self.rnn.init_zero_hidden(batch_size, self.device)
|
| 110 |
-
return hidden_state
|
| 111 |
-
|
| 112 |
-
def forward(self, images, captions=None, hidden_state=None):
|
| 113 |
-
B = images.size(0)
|
| 114 |
-
features = self.encoder(images)
|
| 115 |
-
captions = self._normalize_captions(captions)
|
| 116 |
-
hidden_state = self._ensure_hidden(hidden_state, B)
|
| 117 |
-
if captions is not None:
|
| 118 |
-
inputs = captions[:, :-1]
|
| 119 |
-
logits, _ = self.rnn(inputs, hidden_state, features)
|
| 120 |
-
return logits
|
| 121 |
-
tokens, _ = self.rnn.generate(features, hidden_state, max_len=50)
|
| 122 |
-
return tokens
|
| 123 |
-
|
| 124 |
-
def train_setup(self, prm):
|
| 125 |
-
self.to(self.device)
|
| 126 |
-
lr = float(prm.get('lr', 1e-3)) if isinstance(prm, dict) else 1e-3
|
| 127 |
-
momentum = float(prm.get('momentum', 0.9)) if isinstance(prm, dict) else 0.9
|
| 128 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=lr, momentum=momentum)
|
| 129 |
-
self.criterion = self.criterion.to(self.device)
|
| 130 |
-
|
| 131 |
-
def learn(self, train_data):
|
| 132 |
-
self.train()
|
| 133 |
-
for images, captions in train_data:
|
| 134 |
-
images = images.to(self.device)
|
| 135 |
-
captions = captions.to(self.device)
|
| 136 |
-
caps_for_loss = self._normalize_captions(captions)
|
| 137 |
-
self.optimizer.zero_grad()
|
| 138 |
-
logits = self(images, captions, None)
|
| 139 |
-
T_eff = caps_for_loss.size(1) - 1
|
| 140 |
-
if logits.size(1) != T_eff:
|
| 141 |
-
T_match = min(logits.size(1), T_eff)
|
| 142 |
-
logits = logits[:, :T_match, :]
|
| 143 |
-
caps_for_loss = caps_for_loss[:, :T_match + 1]
|
| 144 |
-
loss = self.criterion(
|
| 145 |
-
logits.contiguous().view(-1, self.vocab_size),
|
| 146 |
-
caps_for_loss[:, 1:].contiguous().view(-1).long()
|
| 147 |
-
)
|
| 148 |
-
loss.backward()
|
| 149 |
-
torch.nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 150 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C5C-ResNetTransformer-83fb6b6bb7c76b742ad0713d29463514.py
DELETED
|
@@ -1,181 +0,0 @@
|
|
| 1 |
-
import math
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.nn.functional as F
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class SEBlock(nn.Module):
|
| 12 |
-
def __init__(self, channels: int, reduction: int = 4):
|
| 13 |
-
super().__init__()
|
| 14 |
-
hidden = max(1, channels // reduction)
|
| 15 |
-
self.avg = nn.AdaptiveAvgPool2d(1)
|
| 16 |
-
self.fc1 = nn.Conv2d(channels, hidden, kernel_size=1, bias=True)
|
| 17 |
-
self.fc2 = nn.Conv2d(hidden, channels, kernel_size=1, bias=True)
|
| 18 |
-
|
| 19 |
-
def forward(self, x):
|
| 20 |
-
s = self.avg(x)
|
| 21 |
-
s = F.relu(self.fc1(s), inplace=True)
|
| 22 |
-
s = torch.sigmoid(self.fc2(s))
|
| 23 |
-
return x * s
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
class InvertedResidual(nn.Module):
|
| 27 |
-
def __init__(self, in_ch: int, out_ch: int, stride: int = 1, expand: int = 3, se_ratio: float = 0.5):
|
| 28 |
-
super().__init__()
|
| 29 |
-
hidden = in_ch * expand
|
| 30 |
-
self.use_res = (stride == 1 and in_ch == out_ch)
|
| 31 |
-
layers = []
|
| 32 |
-
if expand != 1:
|
| 33 |
-
layers += [nn.Conv2d(in_ch, hidden, 1, bias=False), nn.BatchNorm2d(hidden), nn.SiLU(inplace=True)]
|
| 34 |
-
else:
|
| 35 |
-
hidden = in_ch
|
| 36 |
-
layers += [
|
| 37 |
-
nn.Conv2d(hidden, hidden, 3, stride, 1, groups=hidden, bias=False),
|
| 38 |
-
nn.BatchNorm2d(hidden),
|
| 39 |
-
nn.SiLU(inplace=True),
|
| 40 |
-
]
|
| 41 |
-
layers += [nn.Conv2d(hidden, out_ch, 1, bias=False), nn.BatchNorm2d(out_ch)]
|
| 42 |
-
self.block = nn.Sequential(*layers)
|
| 43 |
-
red = max(1, int(round(1.0 / se_ratio))) if se_ratio > 0 else 4
|
| 44 |
-
self.se = SEBlock(out_ch, reduction=red) if se_ratio > 0 else nn.Identity()
|
| 45 |
-
|
| 46 |
-
def forward(self, x):
|
| 47 |
-
out = self.block(x)
|
| 48 |
-
out = self.se(out)
|
| 49 |
-
if self.use_res:
|
| 50 |
-
out = out + x
|
| 51 |
-
return out
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
class Encoder(nn.Module):
|
| 55 |
-
def __init__(self, in_channels: int, hidden_dim: int = 512, se_ratio: float = 0.5):
|
| 56 |
-
super().__init__()
|
| 57 |
-
c1, c2, c3 = 64, 128, hidden_dim
|
| 58 |
-
self.stem = nn.Sequential(
|
| 59 |
-
nn.Conv2d(in_channels, c1, 3, 2, 1, bias=False),
|
| 60 |
-
nn.BatchNorm2d(c1),
|
| 61 |
-
nn.SiLU(inplace=True),
|
| 62 |
-
)
|
| 63 |
-
self.stage1 = InvertedResidual(c1, c1, stride=1, expand=3, se_ratio=se_ratio)
|
| 64 |
-
self.stage2 = InvertedResidual(c1, c2, stride=2, expand=3, se_ratio=se_ratio)
|
| 65 |
-
self.stage3 = InvertedResidual(c2, c3, stride=2, expand=3, se_ratio=se_ratio)
|
| 66 |
-
self.pool = nn.AdaptiveAvgPool2d(1)
|
| 67 |
-
|
| 68 |
-
def forward(self, x):
|
| 69 |
-
x = self.stem(x)
|
| 70 |
-
x = self.stage1(x)
|
| 71 |
-
x = self.stage2(x)
|
| 72 |
-
x = self.stage3(x)
|
| 73 |
-
x = self.pool(x).flatten(1)
|
| 74 |
-
return x.unsqueeze(1)
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
class PositionalEncoding(nn.Module):
|
| 78 |
-
def __init__(self, d_model: int, max_len: int = 5000):
|
| 79 |
-
super().__init__()
|
| 80 |
-
pe = torch.zeros(max_len, d_model)
|
| 81 |
-
pos = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
|
| 82 |
-
div = torch.exp(torch.arange(0, d_model, 2, dtype=torch.float) * (-math.log(10000.0) / d_model))
|
| 83 |
-
pe[:, 0::2] = torch.sin(pos * div)
|
| 84 |
-
pe[:, 1::2] = torch.cos(pos * div)
|
| 85 |
-
pe = pe.unsqueeze(0)
|
| 86 |
-
self.register_buffer("pe", pe, persistent=False)
|
| 87 |
-
|
| 88 |
-
def forward(self, x):
|
| 89 |
-
T = x.size(1)
|
| 90 |
-
return x + self.pe[:, :T, :]
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
class TransformerShim(nn.Module):
|
| 94 |
-
def __init__(self, vocab_size: int, d_model: int = 512, nhead: int = 8, num_layers: int = 1, dim_ff: int = 2048):
|
| 95 |
-
super().__init__()
|
| 96 |
-
assert d_model % nhead == 0
|
| 97 |
-
self.d_model = d_model
|
| 98 |
-
self.embedding = nn.Embedding(vocab_size, d_model)
|
| 99 |
-
self.pos = PositionalEncoding(d_model)
|
| 100 |
-
layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead, dim_feedforward=dim_ff, batch_first=True)
|
| 101 |
-
self.dec = nn.TransformerDecoder(layer, num_layers=num_layers)
|
| 102 |
-
self.fc = nn.Linear(d_model, vocab_size)
|
| 103 |
-
self.num_layers = num_layers
|
| 104 |
-
|
| 105 |
-
def init_zero_hidden(self, batch: int, device: torch.device):
|
| 106 |
-
h0 = torch.zeros(self.num_layers, batch, self.d_model, device=device)
|
| 107 |
-
c0 = torch.zeros(self.num_layers, batch, self.d_model, device=device)
|
| 108 |
-
return (h0, c0)
|
| 109 |
-
|
| 110 |
-
def forward(self, inputs: torch.Tensor, hidden_state, features: torch.Tensor):
|
| 111 |
-
x = self.embedding(inputs) * math.sqrt(self.d_model)
|
| 112 |
-
x = self.pos(x)
|
| 113 |
-
T = inputs.size(1)
|
| 114 |
-
mask = torch.triu(torch.full((T, T), float('-inf'), device=inputs.device), diagonal=1)
|
| 115 |
-
out = self.dec(tgt=x, memory=features, tgt_mask=mask)
|
| 116 |
-
logits = self.fc(out)
|
| 117 |
-
return logits, hidden_state
|
| 118 |
-
|
| 119 |
-
@torch.no_grad()
|
| 120 |
-
def greedy_decode(self, features: torch.Tensor, max_len: int, sos: int = 1, eos: int = 2):
|
| 121 |
-
B = features.size(0)
|
| 122 |
-
ys = torch.full((B, 1), sos, dtype=torch.long, device=features.device)
|
| 123 |
-
for _ in range(max_len):
|
| 124 |
-
x = self.embedding(ys) * math.sqrt(self.d_model)
|
| 125 |
-
x = self.pos(x)
|
| 126 |
-
T = x.size(1)
|
| 127 |
-
mask = torch.triu(torch.full((T, T), float('-inf'), device=ys.device), diagonal=1)
|
| 128 |
-
out = self.dec(tgt=x, memory=features, tgt_mask=mask)
|
| 129 |
-
next_logits = self.fc(out[:, -1, :])
|
| 130 |
-
next_ids = next_logits.argmax(dim=-1, keepdim=True)
|
| 131 |
-
ys = torch.cat([ys, next_ids], dim=1)
|
| 132 |
-
if (next_ids == eos).all():
|
| 133 |
-
break
|
| 134 |
-
return ys[:, 1:]
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
class Net(nn.Module):
|
| 138 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 139 |
-
super().__init__()
|
| 140 |
-
self.device = device
|
| 141 |
-
in_channels = int(in_shape[1])
|
| 142 |
-
vocab_size = int(out_shape[0])
|
| 143 |
-
hidden_dim = int(prm.get('hidden_dim', 512))
|
| 144 |
-
nhead = 8 if hidden_dim % 8 == 0 else 4
|
| 145 |
-
self.encoder = Encoder(in_channels, hidden_dim=hidden_dim, se_ratio=0.5)
|
| 146 |
-
self.rnn = TransformerShim(vocab_size=vocab_size, d_model=hidden_dim, nhead=nhead, num_layers=1, dim_ff=2048)
|
| 147 |
-
self.vocab_size = vocab_size
|
| 148 |
-
|
| 149 |
-
def train_setup(self, prm):
|
| 150 |
-
self.to(self.device)
|
| 151 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=0).to(self.device),)
|
| 152 |
-
beta1 = float(prm.get('momentum', 0.9))
|
| 153 |
-
self.optimizer = torch.optim.AdamW(self.parameters(), lr=float(prm['lr']), betas=(beta1, 0.999), weight_decay=0.01)
|
| 154 |
-
|
| 155 |
-
def learn(self, train_data):
|
| 156 |
-
self.train()
|
| 157 |
-
for images, captions in train_data:
|
| 158 |
-
images = images.to(self.device)
|
| 159 |
-
captions = captions.to(self.device)
|
| 160 |
-
logits, _ = self(images, captions, None)
|
| 161 |
-
tgt = (captions[:, 0, :] if captions.ndim == 3 else captions)[:, 1:]
|
| 162 |
-
loss = self.criteria[0](logits.reshape(-1, logits.size(-1)), tgt.reshape(-1))
|
| 163 |
-
self.optimizer.zero_grad()
|
| 164 |
-
loss.backward()
|
| 165 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3.0)
|
| 166 |
-
self.optimizer.step()
|
| 167 |
-
|
| 168 |
-
def forward(self, images, captions=None, hidden_state=None):
|
| 169 |
-
assert images.dim() == 4
|
| 170 |
-
features = self.encoder(images)
|
| 171 |
-
if captions is None:
|
| 172 |
-
return self.rnn.greedy_decode(features, max_len=50)
|
| 173 |
-
if captions.ndim == 3:
|
| 174 |
-
captions = captions[:, 0, :]
|
| 175 |
-
inputs = captions[:, :-1]
|
| 176 |
-
assert inputs.dtype == torch.long
|
| 177 |
-
if hidden_state is None:
|
| 178 |
-
hidden_state = self.rnn.init_zero_hidden(images.size(0), images.device)
|
| 179 |
-
logits, hidden_state = self.rnn(inputs, hidden_state, features)
|
| 180 |
-
assert logits.dim() == 3 and logits.shape[1] == inputs.shape[1]
|
| 181 |
-
return logits, hidden_state
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/C8C-ResNetTransformer-7730b6eb6979d27e2e1bbc7d05255dff.py
DELETED
|
@@ -1,239 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class ChannelAttention(nn.Module):
|
| 11 |
-
def __init__(self, channel, reduction=4):
|
| 12 |
-
super().__init__()
|
| 13 |
-
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
| 14 |
-
self.max_pool = nn.AdaptiveMaxPool2d(1)
|
| 15 |
-
self.shared_mlp = nn.Sequential(
|
| 16 |
-
nn.Conv2d(channel, channel // reduction, 1, bias=False),
|
| 17 |
-
nn.ReLU(),
|
| 18 |
-
nn.Conv2d(channel // reduction, channel, 1, bias=False)
|
| 19 |
-
)
|
| 20 |
-
self.sigmoid = nn.Sigmoid()
|
| 21 |
-
|
| 22 |
-
def forward(self, x):
|
| 23 |
-
avg_out = self.shared_mlp(self.avg_pool(x))
|
| 24 |
-
max_out = self.shared_mlp(self.max_pool(x))
|
| 25 |
-
return self.sigmoid(avg_out + max_out)
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
class SpatialAttention(nn.Module):
|
| 29 |
-
def __init__(self, kernel_size=7):
|
| 30 |
-
super().__init__()
|
| 31 |
-
self.pool = nn.MaxPool2d(kernel_size, stride=1, padding=kernel_size//2)
|
| 32 |
-
self.sigmoid = nn.Sigmoid()
|
| 33 |
-
|
| 34 |
-
def forward(self, x):
|
| 35 |
-
edge = self.pool(x)
|
| 36 |
-
return x * self.sigmoid(edge)
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
class CABlock(nn.Module):
|
| 40 |
-
def __init__(self, channel, reduction=4, kernel_size=7):
|
| 41 |
-
super().__init__()
|
| 42 |
-
self.channel_att = ChannelAttention(channel, reduction)
|
| 43 |
-
self.spatial_att = SpatialAttention(kernel_size)
|
| 44 |
-
self.conv = nn.Conv2d(channel, channel, 1)
|
| 45 |
-
|
| 46 |
-
def forward(self, x):
|
| 47 |
-
out = x.clone()
|
| 48 |
-
out = self.conv(out)
|
| 49 |
-
out = self.channel_att(out) * out + out
|
| 50 |
-
out = self.spatial_att(out)
|
| 51 |
-
return out
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
class CBAM(nn.Module):
|
| 55 |
-
def __init__(self, channel, reduction=4, kernel_size=7):
|
| 56 |
-
super().__init__()
|
| 57 |
-
self.channel = ChannelAttention(channel, reduction)
|
| 58 |
-
self.spatial = SpatialAttention(kernel_size)
|
| 59 |
-
|
| 60 |
-
def forward(self, x):
|
| 61 |
-
return self.channel(x) * x + self.spatial(x) * x
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
class MBConvBlock(nn.Module):
|
| 65 |
-
def __init__(self, in_channels, out_channels, expand=4.0, kernel_size=5, stride=1, se_ratio=4.0, drop_rate=0.1, index=0):
|
| 66 |
-
super().__init__()
|
| 67 |
-
self.has_se = (out_channels != in_channels)
|
| 68 |
-
self.depth_multiplier = expand
|
| 69 |
-
self.pointwise_conv1 = nn.Conv2d(in_channels, int(in_channels * self.depth_multiplier), 1, bias=False)
|
| 70 |
-
self.bn1 = nn.BatchNorm2d(int(in_channels * self.depth_multiplier))
|
| 71 |
-
self.act1 = nn.GELU()
|
| 72 |
-
self.depth_conv = nn.Conv2d(int(in_channels * self.depth_multiplier), int(in_channels * self.depth_multiplier), kernel_size, padding=kernel_size//2, groups=int(in_channels * self.depth_multiplier), stride=stride, bias=False)
|
| 73 |
-
self.bn2 = nn.BatchNorm2d(int(in_channels * self.depth_multiplier))
|
| 74 |
-
self.act2 = nn.GELU()
|
| 75 |
-
self.pointwise_conv2 = nn.Conv2d(int(in_channels * self.depth_multiplier), out_channels, 1, bias=False)
|
| 76 |
-
self.bn3 = nn.BatchNorm2d(out_channels) if self.has_se else None
|
| 77 |
-
mid_channels = max(1, int(out_channels // (int(se_ratio) if isinstance(se_ratio, (int, float)) and se_ratio >= 1 else 1)))
|
| 78 |
-
self.se = nn.Sequential(
|
| 79 |
-
nn.AdaptiveAvgPool2d(1),
|
| 80 |
-
nn.Flatten(1),
|
| 81 |
-
nn.Linear(int(in_channels * self.depth_multiplier), mid_channels, bias=True),
|
| 82 |
-
nn.ReLU(inplace=True),
|
| 83 |
-
nn.Linear(mid_channels, out_channels, bias=True),
|
| 84 |
-
nn.Sigmoid(),
|
| 85 |
-
nn.Unflatten(1, (out_channels, 1, 1))
|
| 86 |
-
)
|
| 87 |
-
|
| 88 |
-
def forward(self, x):
|
| 89 |
-
x = self.pointwise_conv1(x)
|
| 90 |
-
x = self.bn1(x)
|
| 91 |
-
x = self.act1(x)
|
| 92 |
-
if self.depth_conv.kernel_size == (1, 1):
|
| 93 |
-
shortcut = x
|
| 94 |
-
else:
|
| 95 |
-
x = self.depth_conv(x)
|
| 96 |
-
x = self.bn2(x)
|
| 97 |
-
x = self.act2(x)
|
| 98 |
-
shortcut = None
|
| 99 |
-
x = self.pointwise_conv2(x)
|
| 100 |
-
if self.has_se and shortcut is not None and x.shape[1] == shortcut.shape[1]:
|
| 101 |
-
x = self.se(x + shortcut)
|
| 102 |
-
return x
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
class ScConv(nn.Module):
|
| 106 |
-
def __init__(self, in_channels, out_channels, kernel_size=7, stride=1, groups=64, reduction=4, deploy=False):
|
| 107 |
-
super().__init__()
|
| 108 |
-
self.identity_connection = in_channels == out_channels and stride == 1
|
| 109 |
-
self.padding = kernel_size // 2
|
| 110 |
-
self.empty = False if not deploy else False
|
| 111 |
-
self.pointwise1 = nn.Conv2d(in_channels, out_channels, 1)
|
| 112 |
-
self.depth_conv = nn.Conv2d(in_channels, in_channels, kernel_size, stride, kernel_size, bias=False)
|
| 113 |
-
self.bn2 = nn.BatchNorm2d(in_channels)
|
| 114 |
-
self.act = nn.GELU()
|
| 115 |
-
self.pointwise2 = nn.Conv2d(in_channels, out_channels, 1)
|
| 116 |
-
self.bn3 = nn.BatchNorm2d(out_channels) if self.identity_connection else None
|
| 117 |
-
self.se = nn.Sequential(
|
| 118 |
-
nn.AdaptiveAvgPool2d(1),
|
| 119 |
-
nn.Linear(in_features=in_channels, out_features=in_channels//16)
|
| 120 |
-
) if not self.empty else None
|
| 121 |
-
|
| 122 |
-
def forward(self, x):
|
| 123 |
-
identity = x
|
| 124 |
-
y = self.pointwise1(x)
|
| 125 |
-
if not self.empty:
|
| 126 |
-
y = self.depth_conv(y)
|
| 127 |
-
y = self.bn2(y)
|
| 128 |
-
y = self.act(y)
|
| 129 |
-
y = self.pointwise2(y)
|
| 130 |
-
if self.bn3 is not None:
|
| 131 |
-
y = self.bn3(y)
|
| 132 |
-
if self.identity_connection and y.shape == identity.shape:
|
| 133 |
-
y = identity + y
|
| 134 |
-
return y
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
class Decoder(nn.Module):
|
| 138 |
-
def __init__(self, vocab_size, d_model=768, hidden_size=512):
|
| 139 |
-
super().__init__()
|
| 140 |
-
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
|
| 141 |
-
self.gru = nn.GRU(d_model, hidden_size, batch_first=True)
|
| 142 |
-
self.fc = nn.Linear(hidden_size, vocab_size)
|
| 143 |
-
self.init_from_feat = nn.Linear(d_model, hidden_size)
|
| 144 |
-
|
| 145 |
-
def init_zero_hidden(self, batch, device):
|
| 146 |
-
h0 = torch.zeros(1, batch, self.gru.hidden_size, device=device)
|
| 147 |
-
c0 = torch.zeros_like(h0)
|
| 148 |
-
return (h0, c0)
|
| 149 |
-
|
| 150 |
-
def forward(self, inputs, hidden_state=None, features=None):
|
| 151 |
-
emb = self.embedding(inputs)
|
| 152 |
-
if hidden_state is None:
|
| 153 |
-
if features is not None:
|
| 154 |
-
h = self.init_from_feat(features).unsqueeze(0)
|
| 155 |
-
else:
|
| 156 |
-
h = torch.zeros(1, inputs.size(0), self.gru.hidden_size, device=inputs.device)
|
| 157 |
-
else:
|
| 158 |
-
h = hidden_state[0] if isinstance(hidden_state, tuple) else hidden_state
|
| 159 |
-
out, h = self.gru(emb, h)
|
| 160 |
-
logits = self.fc(out)
|
| 161 |
-
return logits, (h, torch.zeros_like(h))
|
| 162 |
-
|
| 163 |
-
def greedy_decode(self, features, max_len=20, start_id=1, end_id=2):
|
| 164 |
-
B = features.size(0)
|
| 165 |
-
device = features.device
|
| 166 |
-
h = self.init_from_feat(features).unsqueeze(0)
|
| 167 |
-
seq = torch.full((B, 1), start_id, dtype=torch.long, device=device)
|
| 168 |
-
tokens = []
|
| 169 |
-
for _ in range(max_len):
|
| 170 |
-
emb = self.embedding(seq[:, -1:])
|
| 171 |
-
out, h = self.gru(emb, h)
|
| 172 |
-
next_logits = self.fc(out[:, -1, :])
|
| 173 |
-
next_ids = next_logits.argmax(-1, keepdim=True)
|
| 174 |
-
tokens.append(next_ids)
|
| 175 |
-
seq = torch.cat([seq, next_ids], dim=1)
|
| 176 |
-
if (next_ids == end_id).all():
|
| 177 |
-
break
|
| 178 |
-
if len(tokens) == 0:
|
| 179 |
-
return torch.zeros((B, 0), dtype=torch.long, device=device)
|
| 180 |
-
return torch.cat(tokens, dim=1)
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
class Net(nn.Module):
|
| 184 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 185 |
-
super().__init__()
|
| 186 |
-
self.device = device
|
| 187 |
-
vocab_size = int(out_shape[0])
|
| 188 |
-
self.embed_dim = 768
|
| 189 |
-
self.num_heads = 8
|
| 190 |
-
self.num_layers = 6
|
| 191 |
-
self.dropout_rate = 0.1
|
| 192 |
-
self.prj = prm.get('prj', {})
|
| 193 |
-
self.dropout = getattr(prm, 'dropout', 0.1)
|
| 194 |
-
self.attention_dropout = getattr(prm, 'attention_dropout', 0.1)
|
| 195 |
-
self.use_checkpointing = getattr(prm, 'use_checkpointing', False)
|
| 196 |
-
self.use_mem_efficient = getattr(prm, 'use_mem_efficient', True)
|
| 197 |
-
self.projection = nn.Conv2d(3, self.embed_dim, 3, bias=False)
|
| 198 |
-
self.tokenization = nn.AdaptiveAvgPool2d(1)
|
| 199 |
-
self.cls_token = nn.Parameter(torch.randn(1, 1, self.embed_dim))
|
| 200 |
-
self.hybrid_encoder = nn.Sequential(
|
| 201 |
-
CBAM(64),
|
| 202 |
-
CBAM(64),
|
| 203 |
-
ScConv(64, 64),
|
| 204 |
-
MBConvBlock(64, 64)
|
| 205 |
-
)
|
| 206 |
-
self.transformer = nn.TransformerEncoderLayer(d_model=self.embed_dim, nhead=self.num_heads, dropout=self.dropout_rate, batch_first=True)
|
| 207 |
-
self.fc = nn.Linear(self.embed_dim, vocab_size)
|
| 208 |
-
self.rnn = Decoder(vocab_size=vocab_size, d_model=self.embed_dim, hidden_size=640)
|
| 209 |
-
|
| 210 |
-
def train_setup(self, prm):
|
| 211 |
-
self.to(self.device)
|
| 212 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=0).to(self.device),)
|
| 213 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 214 |
-
|
| 215 |
-
def forward(self, images, captions=None, hidden_state=None):
|
| 216 |
-
assert images.dim() == 4
|
| 217 |
-
feats = self.tokenization(self.projection(images)).flatten(1)
|
| 218 |
-
if captions is not None:
|
| 219 |
-
if captions.ndim == 3:
|
| 220 |
-
captions = captions[:, 0, :]
|
| 221 |
-
inputs = captions[:, :-1]
|
| 222 |
-
logits, hidden_state = self.rnn(inputs, hidden_state, features=feats)
|
| 223 |
-
assert logits.shape[1] == inputs.shape[1]
|
| 224 |
-
return logits, hidden_state
|
| 225 |
-
preds = self.rnn.greedy_decode(feats, max_len=20)
|
| 226 |
-
return preds
|
| 227 |
-
|
| 228 |
-
def learn(self, train_data):
|
| 229 |
-
self.train()
|
| 230 |
-
for images, captions in train_data:
|
| 231 |
-
images = images.to(self.device)
|
| 232 |
-
captions = captions.to(self.device)
|
| 233 |
-
logits, _ = self.forward(images, captions, None)
|
| 234 |
-
tgt = (captions[:, 0, :] if captions.ndim == 3 else captions)[:, 1:]
|
| 235 |
-
loss = self.criteria[0](logits.reshape(-1, logits.size(-1)), tgt.reshape(-1))
|
| 236 |
-
self.optimizer.zero_grad()
|
| 237 |
-
loss.backward()
|
| 238 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 239 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ComplexNet.py
DELETED
|
@@ -1,295 +0,0 @@
|
|
| 1 |
-
from collections import OrderedDict
|
| 2 |
-
|
| 3 |
-
import torch
|
| 4 |
-
import torch.nn as nn
|
| 5 |
-
import torch.nn.functional as F
|
| 6 |
-
from torch._C import _disabled_torch_function_impl
|
| 7 |
-
from torch.nn import init, Module, Conv2d, Linear
|
| 8 |
-
from torch.nn.functional import relu, max_pool2d
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
def _retrieve_elements_from_indices(tensor, indices):
|
| 12 |
-
flattened_tensor = tensor.flatten(start_dim=-2)
|
| 13 |
-
output = flattened_tensor.gather(dim=-1, index=indices.flatten(start_dim=-2)).view_as(indices)
|
| 14 |
-
return output
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
def apply_complex(fr, fi, input, dtype=torch.complex64):
|
| 18 |
-
return (fr(input.real) - fi(input.imag)).type(dtype) \
|
| 19 |
-
+ 1j * (fr(input.imag) + fi(input.real)).type(dtype)
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
def complex_relu(input):
|
| 23 |
-
return relu(input.real).type(torch.complex64) + 1j * relu(input.imag).type(torch.complex64)
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
def complex_max_pool2d(input, kernel_size, stride=None, padding=0,
|
| 27 |
-
dilation=1, ceil_mode=False, return_indices=False):
|
| 28 |
-
absolute_value, indices = max_pool2d(
|
| 29 |
-
input.abs(),
|
| 30 |
-
kernel_size=kernel_size,
|
| 31 |
-
stride=stride,
|
| 32 |
-
padding=padding,
|
| 33 |
-
dilation=dilation,
|
| 34 |
-
ceil_mode=ceil_mode,
|
| 35 |
-
return_indices=True
|
| 36 |
-
)
|
| 37 |
-
absolute_value = absolute_value.type(torch.complex64)
|
| 38 |
-
angle = torch.atan2(input.imag, input.real)
|
| 39 |
-
angle = _retrieve_elements_from_indices(angle, indices)
|
| 40 |
-
return absolute_value \
|
| 41 |
-
* (torch.cos(angle).type(torch.complex64) + 1j * torch.sin(angle).type(torch.complex64))
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
class _ParameterMeta(torch._C._TensorMeta):
|
| 45 |
-
def __instancecheck__(self, instance):
|
| 46 |
-
if self is Parameter:
|
| 47 |
-
if isinstance(instance, torch.Tensor) and getattr(
|
| 48 |
-
instance, "_is_param", False
|
| 49 |
-
):
|
| 50 |
-
return True
|
| 51 |
-
return super().__instancecheck__(instance)
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
class Parameter(torch.Tensor, metaclass=_ParameterMeta):
|
| 55 |
-
def __new__(cls, data=None, requires_grad=True):
|
| 56 |
-
if data is None:
|
| 57 |
-
data = torch.empty(0)
|
| 58 |
-
if type(data) is torch.Tensor or type(data) is Parameter:
|
| 59 |
-
return torch.Tensor._make_subclass(cls, data, requires_grad)
|
| 60 |
-
|
| 61 |
-
t = data.detach().requires_grad_(requires_grad)
|
| 62 |
-
if type(t) is not type(data):
|
| 63 |
-
raise RuntimeError(
|
| 64 |
-
f"Creating a Parameter from an instance of type {type(data).__name__} "
|
| 65 |
-
"requires that detach() returns an instance of the same type, but return "
|
| 66 |
-
f"type {type(t).__name__} was found instead. To use the type as a "
|
| 67 |
-
"Parameter, please correct the detach() semantics defined by "
|
| 68 |
-
"its __torch_dispatch__() implementation."
|
| 69 |
-
)
|
| 70 |
-
t._is_param = True
|
| 71 |
-
return t
|
| 72 |
-
|
| 73 |
-
def __deepcopy__(self, memo):
|
| 74 |
-
if id(self) in memo:
|
| 75 |
-
return memo[id(self)]
|
| 76 |
-
else:
|
| 77 |
-
result = type(self)(
|
| 78 |
-
self.data.clone(memory_format=torch.preserve_format), self.requires_grad
|
| 79 |
-
)
|
| 80 |
-
memo[id(self)] = result
|
| 81 |
-
return result
|
| 82 |
-
|
| 83 |
-
def __repr__(self):
|
| 84 |
-
return "Parameter containing:\n" + super().__repr__()
|
| 85 |
-
|
| 86 |
-
def __reduce_ex__(self, proto):
|
| 87 |
-
state = torch._utils._get_obj_state(self)
|
| 88 |
-
|
| 89 |
-
hooks = OrderedDict()
|
| 90 |
-
if not state:
|
| 91 |
-
return (
|
| 92 |
-
torch._utils._rebuild_parameter,
|
| 93 |
-
(self.data, self.requires_grad, hooks),
|
| 94 |
-
)
|
| 95 |
-
|
| 96 |
-
return (
|
| 97 |
-
torch._utils._rebuild_parameter_with_state,
|
| 98 |
-
(self.data, self.requires_grad, hooks, state),
|
| 99 |
-
)
|
| 100 |
-
|
| 101 |
-
__torch_function__ = _disabled_torch_function_impl
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
class _ComplexBatchNorm(Module):
|
| 105 |
-
|
| 106 |
-
def __init__(self, num_features, eps=1e-5, momentum=0.1, affine=True,
|
| 107 |
-
track_running_stats=True):
|
| 108 |
-
super(_ComplexBatchNorm, self).__init__()
|
| 109 |
-
self.num_features = num_features
|
| 110 |
-
self.eps = eps
|
| 111 |
-
self.momentum = momentum
|
| 112 |
-
self.affine = affine
|
| 113 |
-
self.track_running_stats = track_running_stats
|
| 114 |
-
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
| 115 |
-
if self.affine:
|
| 116 |
-
self.weight = Parameter(torch.Tensor(num_features, 3)).to(self.device)
|
| 117 |
-
self.bias = Parameter(torch.Tensor(num_features, 2)).to(self.device)
|
| 118 |
-
else:
|
| 119 |
-
self.register_parameter('weight', None)
|
| 120 |
-
self.register_parameter('bias', None)
|
| 121 |
-
if self.track_running_stats:
|
| 122 |
-
self.register_buffer('running_mean', torch.zeros(num_features, dtype=torch.complex64))
|
| 123 |
-
self.register_buffer('running_covar', torch.zeros(num_features, 3))
|
| 124 |
-
self.running_covar[:, 0] = 1.4142135623730951
|
| 125 |
-
self.running_covar[:, 1] = 1.4142135623730951
|
| 126 |
-
self.register_buffer('num_batches_tracked', torch.tensor(0, dtype=torch.long))
|
| 127 |
-
else:
|
| 128 |
-
self.register_parameter('running_mean', None)
|
| 129 |
-
self.register_parameter('running_covar', None)
|
| 130 |
-
self.register_parameter('num_batches_tracked', None)
|
| 131 |
-
self.reset_parameters()
|
| 132 |
-
|
| 133 |
-
def reset_running_stats(self):
|
| 134 |
-
if self.track_running_stats:
|
| 135 |
-
self.running_mean.zero_()
|
| 136 |
-
self.running_covar.zero_()
|
| 137 |
-
self.running_covar[:, 0] = 1.4142135623730951
|
| 138 |
-
self.running_covar[:, 1] = 1.4142135623730951
|
| 139 |
-
self.num_batches_tracked.zero_()
|
| 140 |
-
|
| 141 |
-
def reset_parameters(self):
|
| 142 |
-
self.reset_running_stats()
|
| 143 |
-
if self.affine:
|
| 144 |
-
init.constant_(self.weight[:, :2], 1.4142135623730951)
|
| 145 |
-
init.zeros_(self.weight[:, 2])
|
| 146 |
-
init.zeros_(self.bias)
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
class ComplexBatchNorm2d(_ComplexBatchNorm):
|
| 150 |
-
|
| 151 |
-
def forward(self, input):
|
| 152 |
-
exponential_average_factor = 0.0
|
| 153 |
-
|
| 154 |
-
if self.training and self.track_running_stats:
|
| 155 |
-
if self.num_batches_tracked is not None:
|
| 156 |
-
self.num_batches_tracked += 1
|
| 157 |
-
if self.momentum is None:
|
| 158 |
-
exponential_average_factor = 1.0 / float(self.num_batches_tracked)
|
| 159 |
-
else:
|
| 160 |
-
exponential_average_factor = self.momentum
|
| 161 |
-
|
| 162 |
-
if self.training or (not self.training and not self.track_running_stats):
|
| 163 |
-
mean_r = input.real.mean([0, 2, 3]).type(torch.complex64)
|
| 164 |
-
mean_i = input.imag.mean([0, 2, 3]).type(torch.complex64)
|
| 165 |
-
mean = mean_r + 1j * mean_i
|
| 166 |
-
else:
|
| 167 |
-
mean = self.running_mean
|
| 168 |
-
|
| 169 |
-
if self.training and self.track_running_stats:
|
| 170 |
-
with torch.no_grad():
|
| 171 |
-
self.running_mean = exponential_average_factor * mean \
|
| 172 |
-
+ (1 - exponential_average_factor) * self.running_mean
|
| 173 |
-
|
| 174 |
-
input = input - mean[None, :, None, None]
|
| 175 |
-
|
| 176 |
-
if self.training or (not self.training and not self.track_running_stats):
|
| 177 |
-
n = input.numel() / input.size(1)
|
| 178 |
-
Crr = 1. / n * input.real.pow(2).sum(dim=[0, 2, 3]) + self.eps
|
| 179 |
-
Cii = 1. / n * input.imag.pow(2).sum(dim=[0, 2, 3]) + self.eps
|
| 180 |
-
Cri = (input.real.mul(input.imag)).mean(dim=[0, 2, 3])
|
| 181 |
-
else:
|
| 182 |
-
Crr = self.running_covar[:, 0] + self.eps
|
| 183 |
-
Cii = self.running_covar[:, 1] + self.eps
|
| 184 |
-
Cri = self.running_covar[:, 2]
|
| 185 |
-
|
| 186 |
-
if self.training and self.track_running_stats:
|
| 187 |
-
with torch.no_grad():
|
| 188 |
-
self.running_covar[:, 0] = exponential_average_factor * Crr * n / (n - 1) \
|
| 189 |
-
+ (1 - exponential_average_factor) * self.running_covar[:, 0]
|
| 190 |
-
|
| 191 |
-
self.running_covar[:, 1] = exponential_average_factor * Cii * n / (n - 1) \
|
| 192 |
-
+ (1 - exponential_average_factor) * self.running_covar[:, 1]
|
| 193 |
-
|
| 194 |
-
self.running_covar[:, 2] = exponential_average_factor * Cri * n / (n - 1) \
|
| 195 |
-
+ (1 - exponential_average_factor) * self.running_covar[:, 2]
|
| 196 |
-
|
| 197 |
-
det = Crr * Cii - Cri.pow(2)
|
| 198 |
-
s = torch.sqrt(det)
|
| 199 |
-
t = torch.sqrt(Cii + Crr + 2 * s)
|
| 200 |
-
inverse_st = 1.0 / (s * t)
|
| 201 |
-
Rrr = (Cii + s) * inverse_st
|
| 202 |
-
Rii = (Crr + s) * inverse_st
|
| 203 |
-
Rri = -Cri * inverse_st
|
| 204 |
-
|
| 205 |
-
input = (Rrr[None, :, None, None] * input.real + Rri[None, :, None, None] * input.imag).type(torch.complex64) \
|
| 206 |
-
+ 1j * (Rii[None, :, None, None] * input.imag + Rri[None, :, None, None] * input.real).type(torch.complex64)
|
| 207 |
-
|
| 208 |
-
if self.affine:
|
| 209 |
-
input = (self.weight[None, :, 0, None, None] * input.real + self.weight[None, :, 2, None, None] * input.imag + \
|
| 210 |
-
self.bias[None, :, 0, None, None]).type(torch.complex64) \
|
| 211 |
-
+ 1j * (self.weight[None, :, 2, None, None] * input.real + self.weight[None, :, 1, None, None] * input.imag + \
|
| 212 |
-
self.bias[None, :, 1, None, None]).type(torch.complex64)
|
| 213 |
-
|
| 214 |
-
return input
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
class ComplexConv2d(Module):
|
| 218 |
-
|
| 219 |
-
def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=0,
|
| 220 |
-
dilation=1, groups=1, bias=True):
|
| 221 |
-
super(ComplexConv2d, self).__init__()
|
| 222 |
-
self.conv_r = Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias)
|
| 223 |
-
self.conv_i = Conv2d(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias)
|
| 224 |
-
|
| 225 |
-
def forward(self, input):
|
| 226 |
-
return apply_complex(self.conv_r, self.conv_i, input)
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
class ComplexLinear(Module):
|
| 230 |
-
|
| 231 |
-
def __init__(self, in_features, out_features):
|
| 232 |
-
super(ComplexLinear, self).__init__()
|
| 233 |
-
self.fc_r = Linear(in_features, out_features)
|
| 234 |
-
self.fc_i = Linear(in_features, out_features)
|
| 235 |
-
|
| 236 |
-
def forward(self, input):
|
| 237 |
-
return apply_complex(self.fc_r, self.fc_i, input)
|
| 238 |
-
|
| 239 |
-
|
| 240 |
-
def supported_hyperparameters():
|
| 241 |
-
return {'lr', 'momentum'}
|
| 242 |
-
|
| 243 |
-
|
| 244 |
-
class Net(nn.Module):
|
| 245 |
-
|
| 246 |
-
def train_setup(self, prm):
|
| 247 |
-
self.to(self.device)
|
| 248 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 249 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 250 |
-
|
| 251 |
-
def learn(self, train_data):
|
| 252 |
-
for inputs, labels in train_data:
|
| 253 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 254 |
-
self.optimizer.zero_grad()
|
| 255 |
-
outputs = self(inputs)
|
| 256 |
-
loss = self.criteria[0](outputs, labels)
|
| 257 |
-
loss.backward()
|
| 258 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 259 |
-
self.optimizer.step()
|
| 260 |
-
|
| 261 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 262 |
-
super(Net, self).__init__()
|
| 263 |
-
self.device = device
|
| 264 |
-
self.in_channels = in_shape[1]
|
| 265 |
-
self.in_height = in_shape[2]
|
| 266 |
-
self.in_width = in_shape[3]
|
| 267 |
-
self.conv1 = ComplexConv2d(self.in_channels, 10, 5, 1)
|
| 268 |
-
self.bn = ComplexBatchNorm2d(10)
|
| 269 |
-
self.conv2 = ComplexConv2d(10, 20, 5, 1)
|
| 270 |
-
self.to(self.device)
|
| 271 |
-
tmp_input = torch.full(in_shape, fill_value=0.1).type(torch.complex64).to(self.device)
|
| 272 |
-
x = self.forward1(tmp_input)
|
| 273 |
-
self.interim_size = int(x.view(-1).size()[0] / in_shape[0])
|
| 274 |
-
self.fc1 = ComplexLinear(self.interim_size, 500)
|
| 275 |
-
self.fc2 = ComplexLinear(500, out_shape[0])
|
| 276 |
-
|
| 277 |
-
def forward1(self, x):
|
| 278 |
-
x = x.view(-1, self.in_channels, self.in_height, self.in_width)
|
| 279 |
-
x = self.conv1(x)
|
| 280 |
-
x = complex_relu(x)
|
| 281 |
-
x = complex_max_pool2d(x, 2, 2)
|
| 282 |
-
x = self.bn(x)
|
| 283 |
-
x = complex_relu(self.conv2(x))
|
| 284 |
-
x = complex_max_pool2d(x, 2, 2)
|
| 285 |
-
return x
|
| 286 |
-
|
| 287 |
-
def forward(self, x):
|
| 288 |
-
x = self.forward1(x)
|
| 289 |
-
x = x.view(-1, self.interim_size)
|
| 290 |
-
x = self.fc1(x)
|
| 291 |
-
x = complex_relu(x)
|
| 292 |
-
x = self.fc2(x)
|
| 293 |
-
x = x.abs()
|
| 294 |
-
x = F.log_softmax(x, dim=1)
|
| 295 |
-
return x
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ConditionalDiffusion.py
DELETED
|
@@ -1,230 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import numpy as np
|
| 5 |
-
import os
|
| 6 |
-
import glob
|
| 7 |
-
from PIL import Image
|
| 8 |
-
import itertools
|
| 9 |
-
|
| 10 |
-
from diffusers import AutoencoderKL, UNet2DConditionModel, DDPMScheduler
|
| 11 |
-
from transformers import AutoTokenizer, AutoModel
|
| 12 |
-
|
| 13 |
-
# Optional import for 8-bit optimizer
|
| 14 |
-
try:
|
| 15 |
-
import bitsandbytes as bnb
|
| 16 |
-
|
| 17 |
-
BITSANDBYTES_AVAILABLE = True
|
| 18 |
-
except ImportError:
|
| 19 |
-
BITSANDBYTES_AVAILABLE = False
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
def supported_hyperparameters():
|
| 23 |
-
"""Returns the hyperparameters supported by this model."""
|
| 24 |
-
return {'lr', 'beta1', 'beta2', 'steps_per_epoch'}
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
class Net(nn.Module):
|
| 28 |
-
"""
|
| 29 |
-
The main Net class that holds the Diffusion components and implements the
|
| 30 |
-
framework's training and evaluation logic.
|
| 31 |
-
"""
|
| 32 |
-
|
| 33 |
-
class TextEncoder(nn.Module):
|
| 34 |
-
def __init__(self, out_size=768):
|
| 35 |
-
super().__init__()
|
| 36 |
-
model_name = "distilbert-base-uncased"
|
| 37 |
-
self.tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 38 |
-
self.text_model = AutoModel.from_pretrained(model_name)
|
| 39 |
-
self.text_linear = nn.Linear(768, out_size)
|
| 40 |
-
for param in self.text_model.parameters():
|
| 41 |
-
param.requires_grad = False
|
| 42 |
-
|
| 43 |
-
def forward(self, text):
|
| 44 |
-
device = self.text_linear.weight.device
|
| 45 |
-
inputs = self.tokenizer(text, return_tensors="pt", padding=True, truncation=True)
|
| 46 |
-
outputs = self.text_model(
|
| 47 |
-
input_ids=inputs.input_ids.to(device),
|
| 48 |
-
attention_mask=inputs.attention_mask.to(device)
|
| 49 |
-
)
|
| 50 |
-
return self.text_linear(outputs.last_hidden_state)
|
| 51 |
-
|
| 52 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 53 |
-
super().__init__()
|
| 54 |
-
self.device = device
|
| 55 |
-
self.prm = prm or {}
|
| 56 |
-
self.epoch_counter = 0
|
| 57 |
-
self.model_name = "CLDiffusion"
|
| 58 |
-
|
| 59 |
-
self.vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse").to(device)
|
| 60 |
-
self.vae.requires_grad_(False)
|
| 61 |
-
|
| 62 |
-
self.text_encoder = self.TextEncoder(out_size=prm.get('cross_attention_dim', 768)).to(device)
|
| 63 |
-
|
| 64 |
-
latent_size = in_shape[2] // 8
|
| 65 |
-
|
| 66 |
-
self.unet = UNet2DConditionModel(
|
| 67 |
-
sample_size=latent_size,
|
| 68 |
-
in_channels=4,
|
| 69 |
-
out_channels=4,
|
| 70 |
-
down_block_types=("DownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D"),
|
| 71 |
-
up_block_types=("UpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"),
|
| 72 |
-
block_out_channels=(128, 256, 512),
|
| 73 |
-
cross_attention_dim=prm.get('cross_attention_dim', 768)
|
| 74 |
-
).to(device)
|
| 75 |
-
|
| 76 |
-
# Enable Memory-Efficient Attention (xFormers) if available
|
| 77 |
-
try:
|
| 78 |
-
self.unet.enable_xformers_memory_efficient_attention()
|
| 79 |
-
print("xFormers memory-efficient attention enabled.")
|
| 80 |
-
except Exception:
|
| 81 |
-
print("xFormers not available. Using standard attention.")
|
| 82 |
-
|
| 83 |
-
# Gradient checkpointing is already enabled, which is great for memory saving.
|
| 84 |
-
self.unet.enable_gradient_checkpointing()
|
| 85 |
-
self.noise_scheduler = DDPMScheduler(num_train_timesteps=1000, beta_schedule="squaredcos_cap_v2")
|
| 86 |
-
|
| 87 |
-
# Setup for Mixed-Precision Training
|
| 88 |
-
self.scaler = torch.cuda.amp.GradScaler()
|
| 89 |
-
|
| 90 |
-
# self.checkpoint_dir = os.path.join("checkpoints", self.model_name)
|
| 91 |
-
# if not os.path.exists(self.checkpoint_dir):
|
| 92 |
-
# os.makedirs(self.checkpoint_dir)
|
| 93 |
-
# self.load_checkpoint()
|
| 94 |
-
|
| 95 |
-
# def load_checkpoint(self):
|
| 96 |
-
# # (omitted for brevity - no changes from previous version)
|
| 97 |
-
# unet_files = glob.glob(os.path.join(self.checkpoint_dir, f'{self.model_name}_unet_epoch_*.pth'))
|
| 98 |
-
# text_encoder_files = glob.glob(os.path.join(self.checkpoint_dir, f'{self.model_name}_text_encoder_epoch_*.pth'))
|
| 99 |
-
#
|
| 100 |
-
# if unet_files and text_encoder_files:
|
| 101 |
-
# latest_unet = max(unet_files, key=os.path.getctime)
|
| 102 |
-
# latest_text_encoder = max(text_encoder_files, key=os.path.getctime)
|
| 103 |
-
# print(f"Loading UNet checkpoint: {latest_unet}")
|
| 104 |
-
# print(f"Loading Text Encoder checkpoint: {latest_text_encoder}")
|
| 105 |
-
# self.unet.load_state_dict(torch.load(latest_unet, map_location=self.device))
|
| 106 |
-
# self.text_encoder.load_state_dict(torch.load(latest_text_encoder, map_location=self.device))
|
| 107 |
-
# try:
|
| 108 |
-
# self.epoch_counter = int(os.path.basename(latest_unet).split('_')[-1].split('.')[0])
|
| 109 |
-
# except (ValueError, IndexError):
|
| 110 |
-
# self.epoch_counter = 0
|
| 111 |
-
# else:
|
| 112 |
-
# print("No checkpoint found, starting from scratch.")
|
| 113 |
-
|
| 114 |
-
def train_setup(self, prm):
|
| 115 |
-
trainable_params = list(self.unet.parameters()) + list(self.text_encoder.text_linear.parameters())
|
| 116 |
-
lr = prm['lr']
|
| 117 |
-
beta1 = prm['beta1']
|
| 118 |
-
beta2 = prm['beta2']
|
| 119 |
-
|
| 120 |
-
# Optional 8-bit Optimizer
|
| 121 |
-
# To use, ensure 'bitsandbytes' is installed and uncomment the following lines.
|
| 122 |
-
# if BITSANDBYTES_AVAILABLE:
|
| 123 |
-
# print("Using 8-bit AdamW optimizer.")
|
| 124 |
-
# self.optimizer = bnb.optim.AdamW8bit(trainable_params, lr=lr, betas=(beta1, 0.999))
|
| 125 |
-
# else:
|
| 126 |
-
# print("Using standard AdamW optimizer.")
|
| 127 |
-
self.optimizer = torch.optim.AdamW(trainable_params, lr=lr, betas=(beta1, beta2))
|
| 128 |
-
|
| 129 |
-
self.criterion = nn.MSELoss()
|
| 130 |
-
# self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(self.optimizer, 'max', patience=50, factor=0.5)
|
| 131 |
-
|
| 132 |
-
def learn(self, train_data):
|
| 133 |
-
self.train()
|
| 134 |
-
total_loss = 0.0
|
| 135 |
-
|
| 136 |
-
if not hasattr(self, 'infinite_data_loader'):
|
| 137 |
-
self.infinite_data_loader = itertools.cycle(train_data)
|
| 138 |
-
|
| 139 |
-
num_steps = int(self.prm['steps_per_epoch'] * 400)
|
| 140 |
-
|
| 141 |
-
if num_steps == 0:
|
| 142 |
-
print("Warning: 'steps_per_epoch' is zero. Skipping training for this epoch.")
|
| 143 |
-
return 0.0
|
| 144 |
-
|
| 145 |
-
for i in range(num_steps):
|
| 146 |
-
batch = next(self.infinite_data_loader)
|
| 147 |
-
images, text_prompts = batch
|
| 148 |
-
self.optimizer.zero_grad()
|
| 149 |
-
|
| 150 |
-
with torch.no_grad():
|
| 151 |
-
latents = self.vae.encode(images.to(self.device)).latent_dist.sample() * 0.18215
|
| 152 |
-
|
| 153 |
-
noise = torch.randn_like(latents)
|
| 154 |
-
timesteps = torch.randint(0, self.noise_scheduler.config.num_train_timesteps, (latents.shape[0],),
|
| 155 |
-
device=self.device)
|
| 156 |
-
noisy_latents = self.noise_scheduler.add_noise(latents, noise, timesteps)
|
| 157 |
-
|
| 158 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 159 |
-
|
| 160 |
-
# --- NEW: Mixed-Precision Training Context ---
|
| 161 |
-
with torch.cuda.amp.autocast():
|
| 162 |
-
noise_pred = self.unet(sample=noisy_latents, timestep=timesteps,
|
| 163 |
-
encoder_hidden_states=text_embeddings).sample
|
| 164 |
-
loss = self.criterion(noise_pred, noise)
|
| 165 |
-
|
| 166 |
-
# Scale loss and update weights ---
|
| 167 |
-
self.scaler.scale(loss).backward()
|
| 168 |
-
self.scaler.step(self.optimizer)
|
| 169 |
-
self.scaler.update()
|
| 170 |
-
|
| 171 |
-
total_loss += loss.item()
|
| 172 |
-
|
| 173 |
-
self.epoch_counter += 1
|
| 174 |
-
|
| 175 |
-
# unet_path = os.path.join(self.checkpoint_dir, f"{self.model_name}_unet_epoch_{self.epoch_counter}.pth")
|
| 176 |
-
# text_encoder_path = os.path.join(self.checkpoint_dir,
|
| 177 |
-
# f"{self.model_name}_text_encoder_epoch_{self.epoch_counter}.pth")
|
| 178 |
-
# torch.save(self.unet.state_dict(), unet_path)
|
| 179 |
-
# torch.save(self.text_encoder.state_dict(), text_encoder_path)
|
| 180 |
-
# print(f"\nCompleted epoch {self.epoch_counter}. Saved checkpoint to {unet_path} and {text_encoder_path}")
|
| 181 |
-
|
| 182 |
-
return total_loss / num_steps
|
| 183 |
-
|
| 184 |
-
@torch.no_grad()
|
| 185 |
-
def generate(self, text_prompts, num_inference_steps=50):
|
| 186 |
-
# (omitted for brevity)
|
| 187 |
-
self.eval()
|
| 188 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 189 |
-
latents = torch.randn((len(text_prompts), self.unet.config.in_channels, self.unet.config.sample_size,
|
| 190 |
-
self.unet.config.sample_size), device=self.device)
|
| 191 |
-
self.noise_scheduler.set_timesteps(num_inference_steps)
|
| 192 |
-
for t in self.noise_scheduler.timesteps:
|
| 193 |
-
noise_pred = self.unet(sample=latents, timestep=t, encoder_hidden_states=text_embeddings).sample
|
| 194 |
-
latents = self.noise_scheduler.step(noise_pred, t, latents).prev_sample
|
| 195 |
-
|
| 196 |
-
latents = 1 / 0.18215 * latents
|
| 197 |
-
images = self.vae.decode(latents).sample
|
| 198 |
-
images = (images / 2 + 0.5).clamp(0, 1)
|
| 199 |
-
images = images.cpu().permute(0, 2, 3, 1).numpy()
|
| 200 |
-
return [Image.fromarray((img * 255).astype(np.uint8)) for img in images]
|
| 201 |
-
|
| 202 |
-
@torch.no_grad()
|
| 203 |
-
def forward(self, images, **kwargs):
|
| 204 |
-
# (omitted for brevity - no changes from previous version)
|
| 205 |
-
batch_size = images.size(0)
|
| 206 |
-
fixed_prompts_for_eval = [
|
| 207 |
-
"a photo of a dog", "a painting of a car", "a smiling person"
|
| 208 |
-
]
|
| 209 |
-
prompts_to_use = [fixed_prompts_for_eval[i % len(fixed_prompts_for_eval)] for i in range(batch_size)]
|
| 210 |
-
|
| 211 |
-
output_dir = os.path.join("output_images", self.model_name)
|
| 212 |
-
if not os.path.exists(output_dir):
|
| 213 |
-
os.makedirs(output_dir)
|
| 214 |
-
|
| 215 |
-
# custom_prompts_to_generate = [
|
| 216 |
-
# "a smiling woman with blond hair",
|
| 217 |
-
# "a man wearing eyeglasses"
|
| 218 |
-
# ]
|
| 219 |
-
# if custom_prompts_to_generate:
|
| 220 |
-
# print(f"\n[Inference] Generating {len(custom_prompts_to_generate)} custom image(s)...")
|
| 221 |
-
# custom_images = self.generate(custom_prompts_to_generate)
|
| 222 |
-
# for i, img in enumerate(custom_images):
|
| 223 |
-
# save_path = os.path.join(output_dir,
|
| 224 |
-
# f"{self.model_name}_output_epoch_{self.epoch_counter}_image_{i + 1}.png")
|
| 225 |
-
# img.save(save_path)
|
| 226 |
-
# print(f"[Inference] Saved custom image to {save_path}")
|
| 227 |
-
|
| 228 |
-
eval_images = self.generate(prompts_to_use)
|
| 229 |
-
|
| 230 |
-
return eval_images, prompts_to_use
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ConditionalGAN.py
DELETED
|
@@ -1,278 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import os
|
| 4 |
-
import torchvision.utils as vutils
|
| 5 |
-
from torch.optim.lr_scheduler import LambdaLR
|
| 6 |
-
# --- MODIFICATION: Added for the new 'generate' method ---
|
| 7 |
-
from torchvision.transforms.functional import to_pil_image
|
| 8 |
-
from torch.nn.utils import spectral_norm
|
| 9 |
-
|
| 10 |
-
try:
|
| 11 |
-
from transformers import CLIPTokenizer
|
| 12 |
-
except ImportError:
|
| 13 |
-
raise ImportError("Please install transformers: pip install transformers")
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
def supported_hyperparameters():
|
| 17 |
-
"""Returns the set of hyperparameters supported by this model."""
|
| 18 |
-
return {'lr', 'beta1'}
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
class Self_Attn(nn.Module):
|
| 22 |
-
""" Self attention Layer"""
|
| 23 |
-
|
| 24 |
-
def __init__(self, in_dim):
|
| 25 |
-
super(Self_Attn, self).__init__()
|
| 26 |
-
self.chanel_in = in_dim
|
| 27 |
-
self.query_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)
|
| 28 |
-
self.key_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)
|
| 29 |
-
self.value_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim, kernel_size=1)
|
| 30 |
-
self.gamma = nn.Parameter(torch.zeros(1))
|
| 31 |
-
self.softmax = nn.Softmax(dim=-1)
|
| 32 |
-
|
| 33 |
-
def forward(self, x):
|
| 34 |
-
m_batchsize, C, width, height = x.size()
|
| 35 |
-
proj_query = self.query_conv(x).view(m_batchsize, -1, width * height).permute(0, 2, 1)
|
| 36 |
-
proj_key = self.key_conv(x).view(m_batchsize, -1, width * height)
|
| 37 |
-
energy = torch.bmm(proj_query, proj_key)
|
| 38 |
-
attention = self.softmax(energy)
|
| 39 |
-
proj_value = self.value_conv(x).view(m_batchsize, -1, width * height)
|
| 40 |
-
out = torch.bmm(proj_value, attention.permute(0, 2, 1))
|
| 41 |
-
out = out.view(m_batchsize, C, width, height)
|
| 42 |
-
out = self.gamma * out + x
|
| 43 |
-
return out
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
class Net(nn.Module):
|
| 47 |
-
class Generator(nn.Module):
|
| 48 |
-
# --- No changes needed in Generator subclass ---
|
| 49 |
-
def __init__(self, noise_dim, embed_dim, hidden_dim, vocab_size, img_channels, feature_maps):
|
| 50 |
-
super().__init__()
|
| 51 |
-
self.embedding = nn.Embedding(vocab_size, embed_dim)
|
| 52 |
-
self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True)
|
| 53 |
-
input_dim = noise_dim + hidden_dim
|
| 54 |
-
self.l1 = nn.Sequential(
|
| 55 |
-
spectral_norm(nn.ConvTranspose2d(input_dim, feature_maps * 16, 4, 1, 0, bias=False)),
|
| 56 |
-
nn.BatchNorm2d(feature_maps * 16), nn.ReLU(True))
|
| 57 |
-
self.l2 = nn.Sequential(
|
| 58 |
-
spectral_norm(nn.ConvTranspose2d(feature_maps * 16, feature_maps * 8, 4, 2, 1, bias=False)),
|
| 59 |
-
nn.BatchNorm2d(feature_maps * 8), nn.ReLU(True))
|
| 60 |
-
self.l3 = nn.Sequential(
|
| 61 |
-
spectral_norm(nn.ConvTranspose2d(feature_maps * 8, feature_maps * 4, 4, 2, 1, bias=False)),
|
| 62 |
-
nn.BatchNorm2d(feature_maps * 4), nn.ReLU(True))
|
| 63 |
-
self.attn1 = Self_Attn(feature_maps * 4)
|
| 64 |
-
self.l4 = nn.Sequential(
|
| 65 |
-
spectral_norm(nn.ConvTranspose2d(feature_maps * 4, feature_maps * 2, 4, 2, 1, bias=False)),
|
| 66 |
-
nn.BatchNorm2d(feature_maps * 2), nn.ReLU(True))
|
| 67 |
-
self.l5 = nn.Sequential(
|
| 68 |
-
spectral_norm(nn.ConvTranspose2d(feature_maps * 2, feature_maps, 4, 2, 1, bias=False)),
|
| 69 |
-
nn.BatchNorm2d(feature_maps), nn.ReLU(True))
|
| 70 |
-
self.attn2 = Self_Attn(feature_maps)
|
| 71 |
-
self.l6 = nn.Sequential(
|
| 72 |
-
spectral_norm(nn.ConvTranspose2d(feature_maps, img_channels, 4, 2, 1, bias=False)),
|
| 73 |
-
nn.Tanh())
|
| 74 |
-
|
| 75 |
-
def forward(self, noise, text_tokens):
|
| 76 |
-
embeddings = self.embedding(text_tokens)
|
| 77 |
-
_, (hidden, _) = self.lstm(embeddings)
|
| 78 |
-
text_conditioning = hidden.squeeze(0)
|
| 79 |
-
x = torch.cat([noise, text_conditioning], dim=1)
|
| 80 |
-
x = self.l1(x.unsqueeze(2).unsqueeze(3))
|
| 81 |
-
x = self.l2(x)
|
| 82 |
-
x = self.l3(x)
|
| 83 |
-
x = self.attn1(x)
|
| 84 |
-
x = self.l4(x)
|
| 85 |
-
x = self.l5(x)
|
| 86 |
-
x = self.attn2(x)
|
| 87 |
-
x = self.l6(x)
|
| 88 |
-
return x
|
| 89 |
-
|
| 90 |
-
class Discriminator(nn.Module):
|
| 91 |
-
# --- No changes needed in Discriminator subclass ---
|
| 92 |
-
def __init__(self, embed_dim, hidden_dim, vocab_size, img_channels, feature_maps):
|
| 93 |
-
super().__init__()
|
| 94 |
-
self.embedding = spectral_norm(nn.Embedding(vocab_size, embed_dim))
|
| 95 |
-
self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True)
|
| 96 |
-
self.image_path = nn.Sequential(
|
| 97 |
-
spectral_norm(nn.Conv2d(img_channels, feature_maps, 4, 2, 1, bias=False)),
|
| 98 |
-
nn.LeakyReLU(0.2, inplace=True),
|
| 99 |
-
spectral_norm(nn.Conv2d(feature_maps, feature_maps * 2, 4, 2, 1, bias=False)),
|
| 100 |
-
nn.LeakyReLU(0.2, inplace=True))
|
| 101 |
-
self.text_path = nn.Sequential(
|
| 102 |
-
spectral_norm(nn.Linear(hidden_dim, feature_maps * 2)), nn.ReLU())
|
| 103 |
-
self.combined_path1 = nn.Sequential(
|
| 104 |
-
spectral_norm(nn.Conv2d(feature_maps * 4, feature_maps * 8, 4, 2, 1, bias=False)),
|
| 105 |
-
nn.LeakyReLU(0.2, inplace=True))
|
| 106 |
-
self.attn = Self_Attn(feature_maps * 8)
|
| 107 |
-
self.combined_path2 = nn.Sequential(
|
| 108 |
-
spectral_norm(nn.Conv2d(feature_maps * 8, feature_maps * 16, 4, 2, 1, bias=False)),
|
| 109 |
-
nn.LeakyReLU(0.2, inplace=True),
|
| 110 |
-
spectral_norm(nn.Conv2d(feature_maps * 16, 1, kernel_size=8, stride=1, padding=0, bias=False)))
|
| 111 |
-
|
| 112 |
-
def forward(self, image, text_tokens):
|
| 113 |
-
image_features = self.image_path(image)
|
| 114 |
-
embeddings = self.embedding(text_tokens)
|
| 115 |
-
_, (hidden, _) = self.lstm(embeddings)
|
| 116 |
-
text_conditioning = hidden.squeeze(0)
|
| 117 |
-
text_features = self.text_path(text_conditioning)
|
| 118 |
-
_, _, H, W = image_features.shape
|
| 119 |
-
text_features_replicated = text_features.unsqueeze(2).unsqueeze(3).expand(-1, -1, H, W)
|
| 120 |
-
combined_features = torch.cat([image_features, text_features_replicated], dim=1)
|
| 121 |
-
x = self.combined_path1(combined_features)
|
| 122 |
-
x = self.attn(x)
|
| 123 |
-
x = self.combined_path2(x)
|
| 124 |
-
return x.view(-1)
|
| 125 |
-
|
| 126 |
-
def __init__(self, shape_a, shape_b, prm: dict, device: torch.device) -> None:
|
| 127 |
-
super().__init__()
|
| 128 |
-
self.device = device
|
| 129 |
-
self.prm = prm
|
| 130 |
-
self.vocab_size = 49408
|
| 131 |
-
img_channels = 3
|
| 132 |
-
self.noise_dim = 100
|
| 133 |
-
embed_dim, hidden_dim, feature_maps = 64, 128, 48
|
| 134 |
-
self.tokenizer = CLIPTokenizer.from_pretrained("openai/clip-vit-base-patch32")
|
| 135 |
-
self.max_length = 16
|
| 136 |
-
self.generator = self.Generator(
|
| 137 |
-
self.noise_dim, embed_dim, hidden_dim, self.vocab_size, img_channels, feature_maps
|
| 138 |
-
).to(device)
|
| 139 |
-
self.discriminator = self.Discriminator(
|
| 140 |
-
embed_dim, hidden_dim, self.vocab_size, img_channels, feature_maps
|
| 141 |
-
).to(device)
|
| 142 |
-
self.r1_penalty_weight = 10.0
|
| 143 |
-
|
| 144 |
-
# --- MODIFICATION: New checkpointing logic ---
|
| 145 |
-
# The model name is derived from the config in Train.py and used to create a unique directory
|
| 146 |
-
model_name = self.__class__.__module__.split('.')[-1]
|
| 147 |
-
self.checkpoint_dir = os.path.join('out', 'checkpoints', model_name)
|
| 148 |
-
os.makedirs(self.checkpoint_dir, exist_ok=True)
|
| 149 |
-
self.best_model_path = os.path.join(self.checkpoint_dir, 'best_model.pth')
|
| 150 |
-
self.best_accuracy = -1.0 # Initialize with a very low value
|
| 151 |
-
|
| 152 |
-
def train_setup(self, prm):
|
| 153 |
-
self.to(self.device)
|
| 154 |
-
lr = float(prm.get('lr', 0.0002))
|
| 155 |
-
beta1 = float(prm.get('beta1', 0.5))
|
| 156 |
-
lr_g, lr_d = lr / 4.0, lr
|
| 157 |
-
self.optimizer_G = torch.optim.Adam(self.generator.parameters(), lr=lr_g, betas=(beta1, 0.999))
|
| 158 |
-
self.optimizer_D = torch.optim.Adam(self.discriminator.parameters(), lr=lr_d, betas=(beta1, 0.999))
|
| 159 |
-
total_epochs = 150
|
| 160 |
-
|
| 161 |
-
def lr_lambda(epoch):
|
| 162 |
-
if epoch < total_epochs / 2:
|
| 163 |
-
return 1.0
|
| 164 |
-
else:
|
| 165 |
-
return 1.0 - (epoch - total_epochs / 2) / (total_epochs / 2)
|
| 166 |
-
|
| 167 |
-
self.scheduler_G = LambdaLR(self.optimizer_G, lr_lambda=lr_lambda)
|
| 168 |
-
self.scheduler_D = LambdaLR(self.optimizer_D, lr_lambda=lr_lambda)
|
| 169 |
-
self.criterion = nn.BCEWithLogitsLoss().to(self.device)
|
| 170 |
-
torch.backends.cudnn.benchmark = True
|
| 171 |
-
|
| 172 |
-
# --- MODIFICATION: Resume from the single best checkpoint if it exists ---
|
| 173 |
-
if os.path.exists(self.best_model_path):
|
| 174 |
-
try:
|
| 175 |
-
print(f"--- Found best model checkpoint at {self.best_model_path}. Resuming training. ---")
|
| 176 |
-
# Loads the state dict for the entire Net module (includes G and D)
|
| 177 |
-
self.load_state_dict(torch.load(self.best_model_path, map_location=self.device))
|
| 178 |
-
except Exception as e:
|
| 179 |
-
print(f"Could not load best model checkpoint, starting from scratch. Error: {e}")
|
| 180 |
-
|
| 181 |
-
def learn(self, train_data, current_epoch=0):
|
| 182 |
-
# --- The main training logic for one epoch remains largely the same ---
|
| 183 |
-
for i, data_batch in enumerate(train_data):
|
| 184 |
-
self.generator.train()
|
| 185 |
-
self.discriminator.train()
|
| 186 |
-
real_images, raw_text_prompts = data_batch
|
| 187 |
-
tokenized_prompts = self.tokenizer(
|
| 188 |
-
list(raw_text_prompts), padding='max_length', truncation=True,
|
| 189 |
-
max_length=self.max_length, return_tensors="pt")
|
| 190 |
-
text_tokens = tokenized_prompts['input_ids'].to(self.device)
|
| 191 |
-
real_images = real_images.to(self.device)
|
| 192 |
-
b_size = real_images.size(0)
|
| 193 |
-
real_target = torch.full((b_size,), 0.9, device=self.device)
|
| 194 |
-
fake_target = torch.full((b_size,), 0.1, device=self.device)
|
| 195 |
-
|
| 196 |
-
for _ in range(2): # Update D twice
|
| 197 |
-
self.optimizer_D.zero_grad()
|
| 198 |
-
real_images.requires_grad = True
|
| 199 |
-
output_real = self.discriminator(real_images, text_tokens)
|
| 200 |
-
loss_d_real = self.criterion(output_real, real_target)
|
| 201 |
-
grad_real = torch.autograd.grad(outputs=output_real.sum(), inputs=real_images, create_graph=True)[0]
|
| 202 |
-
grad_penalty = (grad_real.view(grad_real.size(0), -1).norm(2, dim=1) ** 2).mean()
|
| 203 |
-
r1_penalty = self.r1_penalty_weight / 2 * grad_penalty
|
| 204 |
-
with torch.no_grad():
|
| 205 |
-
noise = torch.randn(b_size, self.noise_dim, device=self.device)
|
| 206 |
-
fake_images = self.generator(noise, text_tokens).detach()
|
| 207 |
-
output_fake = self.discriminator(fake_images, text_tokens)
|
| 208 |
-
loss_d_fake = self.criterion(output_fake, fake_target)
|
| 209 |
-
loss_d = loss_d_real + loss_d_fake + r1_penalty
|
| 210 |
-
loss_d.backward()
|
| 211 |
-
self.optimizer_D.step()
|
| 212 |
-
real_images.requires_grad = False
|
| 213 |
-
|
| 214 |
-
self.optimizer_G.zero_grad() # Update G once
|
| 215 |
-
generator_real_target = torch.full((b_size,), 1.0, device=self.device)
|
| 216 |
-
noise_g = torch.randn(b_size, self.noise_dim, device=self.device)
|
| 217 |
-
fake_images_for_g = self.generator(noise_g, text_tokens)
|
| 218 |
-
output_g = self.discriminator(fake_images_for_g, text_tokens)
|
| 219 |
-
loss_g = self.criterion(output_g, generator_real_target)
|
| 220 |
-
loss_g.backward()
|
| 221 |
-
self.optimizer_G.step()
|
| 222 |
-
|
| 223 |
-
if i % 100 == 0:
|
| 224 |
-
print(
|
| 225 |
-
f'[{current_epoch}][{i}/{len(train_data)}] Loss_D: {loss_d.item():.4f} Loss_G: {loss_g.item():.4f}')
|
| 226 |
-
self.scheduler_G.step()
|
| 227 |
-
self.scheduler_D.step()
|
| 228 |
-
|
| 229 |
-
# --- MODIFICATION: Removed the old periodic checkpointing logic from here ---
|
| 230 |
-
return loss_g.item()
|
| 231 |
-
|
| 232 |
-
# --- NEW METHOD: This is called by Train.py after each evaluation ---
|
| 233 |
-
def save_if_best(self, current_accuracy):
|
| 234 |
-
"""
|
| 235 |
-
Saves the model's state_dict only if the current accuracy is the best seen so far.
|
| 236 |
-
"""
|
| 237 |
-
if current_accuracy > self.best_accuracy:
|
| 238 |
-
self.best_accuracy = current_accuracy
|
| 239 |
-
print(f"--- New best accuracy: {current_accuracy:.4f}. Saving model to {self.best_model_path} ---")
|
| 240 |
-
# Save the entire state dict of the Net module
|
| 241 |
-
torch.save(self.state_dict(), self.best_model_path)
|
| 242 |
-
|
| 243 |
-
# --- MODIFICATION: Hijacked forward pass for evaluation (remains the same) ---
|
| 244 |
-
def forward(self, input_tensor: torch.Tensor, text_prompts: list = None):
|
| 245 |
-
self.generator.eval()
|
| 246 |
-
prompts_to_use = ["a red car on the street"] # Fallback prompt
|
| 247 |
-
if text_prompts is not None:
|
| 248 |
-
valid_prompts = [p for p in text_prompts if isinstance(p, str) and p.strip()]
|
| 249 |
-
if valid_prompts:
|
| 250 |
-
prompts_to_use = valid_prompts
|
| 251 |
-
tokenized_prompts = self.tokenizer(
|
| 252 |
-
prompts_to_use, padding='max_length', truncation=True,
|
| 253 |
-
max_length=self.max_length, return_tensors="pt")['input_ids'].to(self.device)
|
| 254 |
-
noise = torch.randn(len(prompts_to_use), self.noise_dim, device=self.device)
|
| 255 |
-
with torch.no_grad():
|
| 256 |
-
generated_tensors = self.generator(noise, tokenized_prompts)
|
| 257 |
-
generated_tensors = generated_tensors * 0.5 + 0.5
|
| 258 |
-
return (generated_tensors, prompts_to_use)
|
| 259 |
-
|
| 260 |
-
# --- NEW METHOD: This is called by save_results.py for final image generation ---
|
| 261 |
-
def generate(self, text_prompts: list):
|
| 262 |
-
"""
|
| 263 |
-
Generates images from a list of text prompts and returns them as PIL Images.
|
| 264 |
-
"""
|
| 265 |
-
self.generator.eval()
|
| 266 |
-
if not text_prompts:
|
| 267 |
-
return []
|
| 268 |
-
tokenized_prompts = self.tokenizer(
|
| 269 |
-
text_prompts, padding='max_length', truncation=True,
|
| 270 |
-
max_length=self.max_length, return_tensors="pt")['input_ids'].to(self.device)
|
| 271 |
-
noise = torch.randn(len(text_prompts), self.noise_dim, device=self.device)
|
| 272 |
-
with torch.no_grad():
|
| 273 |
-
generated_tensors = self.generator(noise, tokenized_prompts)
|
| 274 |
-
generated_tensors = (generated_tensors * 0.5 + 0.5).clamp(0, 1) # Denormalize
|
| 275 |
-
|
| 276 |
-
# Convert tensors to a list of the PIL
|
| 277 |
-
pil_images = [to_pil_image(tensor) for tensor in generated_tensors]
|
| 278 |
-
return pil_images
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ConditionalVAE3.py
DELETED
|
@@ -1,213 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torchvision.models as models
|
| 5 |
-
import torchvision.transforms as T
|
| 6 |
-
import math
|
| 7 |
-
import os
|
| 8 |
-
# Import the required function for saving weights from the framework's utility file.
|
| 9 |
-
from ab.nn.util.Util import export_torch_weights
|
| 10 |
-
from transformers import CLIPTextModel, CLIPTokenizer
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
def supported_hyperparameters():
|
| 14 |
-
"""Returns the hyperparameters supported by this model."""
|
| 15 |
-
# 'save_weights' flag to make checkpointing controllable.
|
| 16 |
-
return {'lr', 'momentum', 'version', 'save_weights'}
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
class PerceptualLoss(nn.Module):
|
| 20 |
-
def __init__(self):
|
| 21 |
-
super(PerceptualLoss, self).__init__()
|
| 22 |
-
vgg = models.vgg16(weights=models.VGG16_Weights.IMAGENET1K_V1).features[:23].eval()
|
| 23 |
-
self.vgg = nn.Sequential(*vgg)
|
| 24 |
-
for param in self.vgg.parameters():
|
| 25 |
-
param.requires_grad = False
|
| 26 |
-
self.l1 = nn.L1Loss()
|
| 27 |
-
self.normalize = T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
|
| 28 |
-
|
| 29 |
-
def forward(self, y_hat, y):
|
| 30 |
-
y_hat_norm = self.normalize(y_hat)
|
| 31 |
-
y_norm = self.normalize(y)
|
| 32 |
-
vgg_y_hat = self.vgg(y_hat_norm)
|
| 33 |
-
vgg_y = self.vgg(y_norm)
|
| 34 |
-
return self.l1(vgg_y_hat, vgg_y)
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
class SelfAttention(nn.Module):
|
| 38 |
-
def __init__(self, in_channels):
|
| 39 |
-
super().__init__()
|
| 40 |
-
self.query = nn.Conv2d(in_channels, in_channels // 8, 1)
|
| 41 |
-
self.key = nn.Conv2d(in_channels, in_channels // 8, 1)
|
| 42 |
-
self.value = nn.Conv2d(in_channels, in_channels, 1)
|
| 43 |
-
self.gamma = nn.Parameter(torch.tensor(0.0))
|
| 44 |
-
|
| 45 |
-
def forward(self, x):
|
| 46 |
-
batch_size, C, width, height = x.size()
|
| 47 |
-
query = self.query(x).view(batch_size, -1, width * height).permute(0, 2, 1)
|
| 48 |
-
key = self.key(x).view(batch_size, -1, width * height)
|
| 49 |
-
attention = torch.bmm(query, key).softmax(dim=-1)
|
| 50 |
-
value = self.value(x).view(batch_size, -1, width * height)
|
| 51 |
-
out = torch.bmm(value, attention.permute(0, 2, 1))
|
| 52 |
-
out = out.view(batch_size, C, width, height)
|
| 53 |
-
return self.gamma * out + x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
class TextEncoder(nn.Module):
|
| 58 |
-
def __init__(self, out_size=128):
|
| 59 |
-
super().__init__()
|
| 60 |
-
model_name = "openai/clip-vit-base-patch32"
|
| 61 |
-
self.tokenizer = CLIPTokenizer.from_pretrained(model_name)
|
| 62 |
-
self.text_model = CLIPTextModel.from_pretrained(model_name)
|
| 63 |
-
self.text_linear = nn.Linear(512, out_size)
|
| 64 |
-
for param in self.text_model.parameters():
|
| 65 |
-
param.requires_grad = False
|
| 66 |
-
|
| 67 |
-
def forward(self, text):
|
| 68 |
-
device = self.text_linear.weight.device
|
| 69 |
-
inputs = self.tokenizer(text, return_tensors="pt", padding=True, truncation=True)
|
| 70 |
-
outputs = self.text_model(
|
| 71 |
-
input_ids=inputs.input_ids.to(device),
|
| 72 |
-
attention_mask=inputs.attention_mask.to(device)
|
| 73 |
-
)
|
| 74 |
-
return self.text_linear(outputs.pooler_output)
|
| 75 |
-
|
| 76 |
-
class CVAE(nn.Module):
|
| 77 |
-
class UpsampleBlock(nn.Module):
|
| 78 |
-
def __init__(self, in_channels, out_channels):
|
| 79 |
-
super().__init__()
|
| 80 |
-
self.conv = nn.Conv2d(in_channels, out_channels * 4, kernel_size=3, padding=1)
|
| 81 |
-
self.pixel_shuffle = nn.PixelShuffle(2)
|
| 82 |
-
self.lrelu = nn.LeakyReLU(0.2, inplace=True)
|
| 83 |
-
|
| 84 |
-
def forward(self, x):
|
| 85 |
-
x = self.conv(x)
|
| 86 |
-
x = self.pixel_shuffle(x)
|
| 87 |
-
x = self.lrelu(x)
|
| 88 |
-
return x
|
| 89 |
-
|
| 90 |
-
def __init__(self, latent_dim=512, text_embedding_dim=128, image_channels=3, image_size=256):
|
| 91 |
-
super().__init__()
|
| 92 |
-
self.latent_dim = latent_dim
|
| 93 |
-
self.encoder_conv = nn.Sequential(
|
| 94 |
-
nn.Conv2d(image_channels, 32, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 95 |
-
nn.Conv2d(32, 64, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 96 |
-
nn.Conv2d(64, 128, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 97 |
-
nn.Conv2d(128, 256, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 98 |
-
nn.Conv2d(256, 512, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 99 |
-
nn.Conv2d(512, 512, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True)
|
| 100 |
-
)
|
| 101 |
-
|
| 102 |
-
with torch.no_grad():
|
| 103 |
-
dummy_input = torch.zeros(1, image_channels, image_size, image_size)
|
| 104 |
-
dummy_output = self.encoder_conv(dummy_input)
|
| 105 |
-
self.final_feature_dim = dummy_output.view(-1).shape[0]
|
| 106 |
-
self.final_conv_shape = dummy_output.shape
|
| 107 |
-
|
| 108 |
-
combined_dim = self.final_feature_dim + text_embedding_dim
|
| 109 |
-
self.fc_mu = nn.Linear(combined_dim, latent_dim)
|
| 110 |
-
self.fc_log_var = nn.Linear(combined_dim, latent_dim)
|
| 111 |
-
self.decoder_input = nn.Linear(latent_dim + text_embedding_dim, self.final_feature_dim)
|
| 112 |
-
|
| 113 |
-
self.decoder_conv = nn.Sequential(
|
| 114 |
-
self.UpsampleBlock(512, 512),
|
| 115 |
-
self.UpsampleBlock(512, 256),
|
| 116 |
-
self.UpsampleBlock(256, 128),
|
| 117 |
-
SelfAttention(128),
|
| 118 |
-
self.UpsampleBlock(128, 64),
|
| 119 |
-
self.UpsampleBlock(64, 32),
|
| 120 |
-
self.UpsampleBlock(32, 16),
|
| 121 |
-
nn.Conv2d(16, image_channels, kernel_size=3, padding=1),
|
| 122 |
-
nn.Tanh()
|
| 123 |
-
)
|
| 124 |
-
|
| 125 |
-
def encode(self, image, text_embedding):
|
| 126 |
-
x = self.encoder_conv(image)
|
| 127 |
-
x = torch.flatten(x, start_dim=1)
|
| 128 |
-
combined = torch.cat([x, text_embedding], dim=1)
|
| 129 |
-
return self.fc_mu(combined), self.fc_log_var(combined)
|
| 130 |
-
|
| 131 |
-
def reparameterize(self, mu, log_var):
|
| 132 |
-
std = torch.exp(0.5 * log_var)
|
| 133 |
-
eps = torch.randn_like(std)
|
| 134 |
-
return mu + eps * std
|
| 135 |
-
|
| 136 |
-
def decode(self, z, text_embedding):
|
| 137 |
-
combined = torch.cat([z, text_embedding], dim=1)
|
| 138 |
-
x = self.decoder_input(combined)
|
| 139 |
-
x = x.view(-1, *self.final_conv_shape[1:])
|
| 140 |
-
return self.decoder_conv(x)
|
| 141 |
-
|
| 142 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 143 |
-
super().__init__()
|
| 144 |
-
self.device = device
|
| 145 |
-
self.prm = prm or {}
|
| 146 |
-
self.text_embedding_dim = 128
|
| 147 |
-
self.latent_dim = 512
|
| 148 |
-
self.model_name = "ConditionalVAE3"
|
| 149 |
-
self.register_buffer('epoch_counter', torch.tensor(0))
|
| 150 |
-
image_channels, image_size = in_shape[1], in_shape[2]
|
| 151 |
-
self.text_encoder = self.TextEncoder(out_size=self.text_embedding_dim).to(device)
|
| 152 |
-
self.cvae = self.CVAE(self.latent_dim, self.text_embedding_dim, image_channels, image_size).to(device)
|
| 153 |
-
|
| 154 |
-
lr = self.prm.get('lr', 1e-4)
|
| 155 |
-
beta1 = self.prm.get('momentum', 0.9)
|
| 156 |
-
self.optimizer = torch.optim.Adam(self.cvae.parameters(), lr=lr, betas=(beta1, 0.999))
|
| 157 |
-
self.reconstruction_loss = nn.L1Loss()
|
| 158 |
-
self.perceptual_loss = PerceptualLoss().to(device)
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
def train_setup(self, prm):
|
| 162 |
-
pass
|
| 163 |
-
|
| 164 |
-
def learn(self, train_data, current_epoch=0):
|
| 165 |
-
self.train()
|
| 166 |
-
self.epoch_counter = torch.tensor(current_epoch)
|
| 167 |
-
total_loss = 0.0
|
| 168 |
-
kld_warmup_epochs = 25
|
| 169 |
-
max_kld_weight = 0.0000025
|
| 170 |
-
|
| 171 |
-
if current_epoch < kld_warmup_epochs:
|
| 172 |
-
kld_weight = max_kld_weight * ((current_epoch + 1) / kld_warmup_epochs)
|
| 173 |
-
else:
|
| 174 |
-
kld_weight = max_kld_weight
|
| 175 |
-
|
| 176 |
-
for batch in train_data:
|
| 177 |
-
real_images, text_prompts = batch
|
| 178 |
-
real_images = real_images.to(self.device)
|
| 179 |
-
self.optimizer.zero_grad()
|
| 180 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 181 |
-
mu, log_var = self.cvae.encode(real_images, text_embeddings)
|
| 182 |
-
z = self.cvae.reparameterize(mu, log_var)
|
| 183 |
-
reconstructed_images = self.cvae.decode(z, text_embeddings)
|
| 184 |
-
recon_loss = self.reconstruction_loss(reconstructed_images, real_images)
|
| 185 |
-
perc_loss = self.perceptual_loss(reconstructed_images, real_images)
|
| 186 |
-
kld_loss = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())
|
| 187 |
-
loss = recon_loss + 0.8 * perc_loss + (kld_weight * kld_loss)
|
| 188 |
-
loss.backward()
|
| 189 |
-
torch.nn.utils.clip_grad_norm_(self.cvae.parameters(), 1.0)
|
| 190 |
-
self.optimizer.step()
|
| 191 |
-
total_loss += loss.item()
|
| 192 |
-
return total_loss / len(train_data) if train_data else 0.0
|
| 193 |
-
|
| 194 |
-
@torch.no_grad()
|
| 195 |
-
def generate(self, text_prompts):
|
| 196 |
-
self.eval()
|
| 197 |
-
num_images = len(text_prompts)
|
| 198 |
-
z = torch.randn(num_images, self.latent_dim, device=self.device)
|
| 199 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 200 |
-
generated_images = self.cvae.decode(z, text_embeddings)
|
| 201 |
-
generated_images = (generated_images + 1) / 2
|
| 202 |
-
return [T.ToPILImage()(img.cpu()) for img in generated_images]
|
| 203 |
-
|
| 204 |
-
@torch.no_grad()
|
| 205 |
-
def forward(self, images, **kwargs):
|
| 206 |
-
prompts_to_use = kwargs.get('prompts')
|
| 207 |
-
if not prompts_to_use:
|
| 208 |
-
batch_size = images.size(0)
|
| 209 |
-
default_prompts = ["a photo of a car"]
|
| 210 |
-
prompts_to_use = [default_prompts[i % len(default_prompts)] for i in range(batch_size)]
|
| 211 |
-
|
| 212 |
-
generated_images = self.generate(prompts_to_use)
|
| 213 |
-
return generated_images, prompts_to_use
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ConditionalVAE4.py
DELETED
|
@@ -1,268 +0,0 @@
|
|
| 1 |
-
# File: ConditionalVAE4.py
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torchvision.models as models
|
| 5 |
-
import torchvision.transforms as T
|
| 6 |
-
import math
|
| 7 |
-
import os
|
| 8 |
-
|
| 9 |
-
from ab.nn.util.Util import export_torch_weights
|
| 10 |
-
from transformers import CLIPTextModel, CLIPTokenizer
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
def supported_hyperparameters():
|
| 14 |
-
"""Returns the hyperparameters supported by this model."""
|
| 15 |
-
#'save_weights' flag to make checkpointing controllable.
|
| 16 |
-
return {'lr', 'momentum', 'version', 'lr_g', 'lr_d', 'save_weights'}
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
class PerceptualLoss(nn.Module):
|
| 20 |
-
def __init__(self):
|
| 21 |
-
super(PerceptualLoss, self).__init__()
|
| 22 |
-
vgg = models.vgg16(weights=models.VGG16_Weights.IMAGENET1K_V1).features[:23].eval()
|
| 23 |
-
self.vgg = nn.Sequential(*vgg)
|
| 24 |
-
for param in self.vgg.parameters():
|
| 25 |
-
param.requires_grad = False
|
| 26 |
-
self.l1 = nn.L1Loss()
|
| 27 |
-
self.normalize = T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
|
| 28 |
-
|
| 29 |
-
def forward(self, y_hat, y):
|
| 30 |
-
y_hat_norm = self.normalize(y_hat)
|
| 31 |
-
y_norm = self.normalize(y)
|
| 32 |
-
vgg_y_hat = self.vgg(y_hat_norm)
|
| 33 |
-
vgg_y = self.vgg(y_norm)
|
| 34 |
-
return self.l1(vgg_y_hat, vgg_y)
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
class SelfAttention(nn.Module):
|
| 38 |
-
def __init__(self, in_channels):
|
| 39 |
-
super().__init__()
|
| 40 |
-
self.query = nn.Conv2d(in_channels, in_channels // 8, 1)
|
| 41 |
-
self.key = nn.Conv2d(in_channels, in_channels // 8, 1)
|
| 42 |
-
self.value = nn.Conv2d(in_channels, in_channels, 1)
|
| 43 |
-
self.gamma = nn.Parameter(torch.tensor(0.0))
|
| 44 |
-
|
| 45 |
-
def forward(self, x):
|
| 46 |
-
batch_size, C, width, height = x.size()
|
| 47 |
-
query = self.query(x).view(batch_size, -1, width * height).permute(0, 2, 1)
|
| 48 |
-
key = self.key(x).view(batch_size, -1, width * height)
|
| 49 |
-
attention = torch.bmm(query, key).softmax(dim=-1)
|
| 50 |
-
value = self.value(x).view(batch_size, -1, width * height)
|
| 51 |
-
out = torch.bmm(value, attention.permute(0, 2, 1))
|
| 52 |
-
out = out.view(batch_size, C, width, height)
|
| 53 |
-
return self.gamma * out + x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
class TextEncoder(nn.Module):
|
| 58 |
-
def __init__(self, out_size=128):
|
| 59 |
-
super().__init__()
|
| 60 |
-
model_name = "openai/clip-vit-base-patch32"
|
| 61 |
-
self.tokenizer = CLIPTokenizer.from_pretrained(model_name)
|
| 62 |
-
self.text_model = CLIPTextModel.from_pretrained(model_name)
|
| 63 |
-
self.text_linear = nn.Linear(512, out_size)
|
| 64 |
-
for param in self.text_model.parameters():
|
| 65 |
-
param.requires_grad = False
|
| 66 |
-
|
| 67 |
-
def forward(self, text):
|
| 68 |
-
device = self.text_linear.weight.device
|
| 69 |
-
inputs = self.tokenizer(text, return_tensors="pt", padding=True, truncation=True)
|
| 70 |
-
outputs = self.text_model(
|
| 71 |
-
input_ids=inputs.input_ids.to(device),
|
| 72 |
-
attention_mask=inputs.attention_mask.to(device)
|
| 73 |
-
)
|
| 74 |
-
return self.text_linear(outputs.pooler_output)
|
| 75 |
-
|
| 76 |
-
class CVAE(nn.Module):
|
| 77 |
-
class UpsampleBlock(nn.Module):
|
| 78 |
-
def __init__(self, in_channels, out_channels):
|
| 79 |
-
super().__init__()
|
| 80 |
-
self.conv = nn.Conv2d(in_channels, out_channels * 4, kernel_size=3, padding=1)
|
| 81 |
-
self.pixel_shuffle = nn.PixelShuffle(2)
|
| 82 |
-
self.lrelu = nn.LeakyReLU(0.2, inplace=True)
|
| 83 |
-
|
| 84 |
-
def forward(self, x):
|
| 85 |
-
x = self.conv(x)
|
| 86 |
-
x = self.pixel_shuffle(x)
|
| 87 |
-
x = self.lrelu(x)
|
| 88 |
-
return x
|
| 89 |
-
|
| 90 |
-
def __init__(self, latent_dim=512, text_embedding_dim=128, image_channels=3, image_size=256):
|
| 91 |
-
super().__init__()
|
| 92 |
-
self.latent_dim = latent_dim
|
| 93 |
-
self.encoder_conv = nn.Sequential(
|
| 94 |
-
nn.Conv2d(image_channels, 32, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 95 |
-
nn.Conv2d(32, 64, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 96 |
-
nn.Conv2d(64, 128, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 97 |
-
nn.Conv2d(128, 256, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 98 |
-
nn.Conv2d(256, 512, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 99 |
-
nn.Conv2d(512, 512, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True)
|
| 100 |
-
)
|
| 101 |
-
|
| 102 |
-
with torch.no_grad():
|
| 103 |
-
dummy_input = torch.zeros(1, image_channels, image_size, image_size)
|
| 104 |
-
dummy_output = self.encoder_conv(dummy_input)
|
| 105 |
-
self.final_feature_dim = dummy_output.view(-1).shape[0]
|
| 106 |
-
self.final_conv_shape = dummy_output.shape
|
| 107 |
-
|
| 108 |
-
combined_dim = self.final_feature_dim + text_embedding_dim
|
| 109 |
-
self.fc_mu = nn.Linear(combined_dim, latent_dim)
|
| 110 |
-
self.fc_log_var = nn.Linear(combined_dim, latent_dim)
|
| 111 |
-
self.decoder_input = nn.Linear(latent_dim + text_embedding_dim, self.final_feature_dim)
|
| 112 |
-
|
| 113 |
-
self.decoder_conv = nn.Sequential(
|
| 114 |
-
self.UpsampleBlock(512, 512),
|
| 115 |
-
self.UpsampleBlock(512, 256),
|
| 116 |
-
self.UpsampleBlock(256, 128),
|
| 117 |
-
SelfAttention(128),
|
| 118 |
-
self.UpsampleBlock(128, 64),
|
| 119 |
-
self.UpsampleBlock(64, 32),
|
| 120 |
-
self.UpsampleBlock(32, 16),
|
| 121 |
-
nn.Conv2d(16, image_channels, kernel_size=3, padding=1),
|
| 122 |
-
nn.Tanh()
|
| 123 |
-
)
|
| 124 |
-
|
| 125 |
-
def encode(self, image, text_embedding):
|
| 126 |
-
x = self.encoder_conv(image)
|
| 127 |
-
x = torch.flatten(x, start_dim=1)
|
| 128 |
-
combined = torch.cat([x, text_embedding], dim=1)
|
| 129 |
-
return self.fc_mu(combined), self.fc_log_var(combined)
|
| 130 |
-
|
| 131 |
-
def reparameterize(self, mu, log_var):
|
| 132 |
-
std = torch.exp(0.5 * log_var)
|
| 133 |
-
eps = torch.randn_like(std)
|
| 134 |
-
return mu + eps * std
|
| 135 |
-
|
| 136 |
-
def decode(self, z, text_embedding):
|
| 137 |
-
combined = torch.cat([z, text_embedding], dim=1)
|
| 138 |
-
x = self.decoder_input(combined)
|
| 139 |
-
x = x.view(-1, *self.final_conv_shape[1:])
|
| 140 |
-
return self.decoder_conv(x)
|
| 141 |
-
|
| 142 |
-
class Discriminator(nn.Module):
|
| 143 |
-
def __init__(self, image_channels=3):
|
| 144 |
-
super().__init__()
|
| 145 |
-
self.model = nn.Sequential(
|
| 146 |
-
nn.Conv2d(image_channels, 64, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 147 |
-
nn.Conv2d(64, 128, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 148 |
-
nn.Conv2d(128, 256, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True),
|
| 149 |
-
nn.Conv2d(256, 1, 4, 1, 0)
|
| 150 |
-
)
|
| 151 |
-
|
| 152 |
-
def forward(self, x):
|
| 153 |
-
return self.model(x)
|
| 154 |
-
|
| 155 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 156 |
-
super().__init__()
|
| 157 |
-
self.device = device
|
| 158 |
-
self.prm = prm or {}
|
| 159 |
-
self.text_embedding_dim = 128
|
| 160 |
-
self.latent_dim = 512
|
| 161 |
-
self.model_name = "ConditionalVAE4"
|
| 162 |
-
|
| 163 |
-
self.register_buffer('epoch_counter', torch.tensor(0))
|
| 164 |
-
|
| 165 |
-
image_channels, image_size = in_shape[1], in_shape[2]
|
| 166 |
-
self.text_encoder = self.TextEncoder(out_size=self.text_embedding_dim).to(device)
|
| 167 |
-
self.cvae = self.CVAE(self.latent_dim, self.text_embedding_dim, image_channels, image_size).to(device)
|
| 168 |
-
self.discriminator = self.Discriminator(image_channels).to(device)
|
| 169 |
-
|
| 170 |
-
lr_g = self.prm.get('lr_g', 2e-6)
|
| 171 |
-
lr_d = self.prm.get('lr_d', 2e-7)
|
| 172 |
-
beta1 = self.prm.get('momentum', 0.5)
|
| 173 |
-
|
| 174 |
-
self.optimizer_g = torch.optim.Adam(self.cvae.parameters(), lr=lr_g, betas=(beta1, 0.999))
|
| 175 |
-
self.optimizer_d = torch.optim.Adam(self.discriminator.parameters(), lr=lr_d, betas=(beta1, 0.999))
|
| 176 |
-
|
| 177 |
-
self.reconstruction_loss = nn.L1Loss()
|
| 178 |
-
self.perceptual_loss = PerceptualLoss().to(device)
|
| 179 |
-
self.adversarial_loss = nn.BCEWithLogitsLoss()
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
def train_setup(self, prm):
|
| 183 |
-
pass
|
| 184 |
-
|
| 185 |
-
def learn(self, train_data, current_epoch=0):
|
| 186 |
-
self.train()
|
| 187 |
-
self.epoch_counter = torch.tensor(current_epoch)
|
| 188 |
-
total_g_loss = 0.0
|
| 189 |
-
total_d_loss = 0.0
|
| 190 |
-
|
| 191 |
-
recon_weight = 10.0
|
| 192 |
-
perc_weight = 1.0
|
| 193 |
-
kld_weight = 0.0000025
|
| 194 |
-
adversarial_weight = 0.0001
|
| 195 |
-
|
| 196 |
-
for batch in train_data:
|
| 197 |
-
real_images, text_prompts = batch
|
| 198 |
-
real_images = real_images.to(self.device)
|
| 199 |
-
|
| 200 |
-
# Train the Discriminator
|
| 201 |
-
self.optimizer_d.zero_grad()
|
| 202 |
-
|
| 203 |
-
with torch.no_grad():
|
| 204 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 205 |
-
mu, log_var = self.cvae.encode(real_images, text_embeddings)
|
| 206 |
-
z = self.cvae.reparameterize(mu, log_var)
|
| 207 |
-
reconstructed_images = self.cvae.decode(z, text_embeddings)
|
| 208 |
-
|
| 209 |
-
real_output = self.discriminator(real_images)
|
| 210 |
-
real_labels = torch.ones_like(real_output, device=self.device)
|
| 211 |
-
fake_labels = torch.zeros_like(real_output, device=self.device)
|
| 212 |
-
|
| 213 |
-
d_loss_real = self.adversarial_loss(real_output, real_labels)
|
| 214 |
-
fake_output = self.discriminator(reconstructed_images.detach())
|
| 215 |
-
d_loss_fake = self.adversarial_loss(fake_output, fake_labels)
|
| 216 |
-
|
| 217 |
-
d_loss = (d_loss_real + d_loss_fake) / 2
|
| 218 |
-
d_loss.backward()
|
| 219 |
-
self.optimizer_d.step()
|
| 220 |
-
total_d_loss += d_loss.item()
|
| 221 |
-
|
| 222 |
-
# Train the VAE (Generator)
|
| 223 |
-
self.optimizer_g.zero_grad()
|
| 224 |
-
|
| 225 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 226 |
-
mu, log_var = self.cvae.encode(real_images, text_embeddings)
|
| 227 |
-
z = self.cvae.reparameterize(mu, log_var)
|
| 228 |
-
reconstructed_images_for_g = self.cvae.decode(z, text_embeddings)
|
| 229 |
-
|
| 230 |
-
recon_loss = self.reconstruction_loss(reconstructed_images_for_g, real_images)
|
| 231 |
-
perc_loss = self.perceptual_loss(reconstructed_images_for_g, real_images)
|
| 232 |
-
kld_loss = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())
|
| 233 |
-
|
| 234 |
-
fake_output_for_g = self.discriminator(reconstructed_images_for_g)
|
| 235 |
-
g_loss_adv = self.adversarial_loss(fake_output_for_g, real_labels)
|
| 236 |
-
|
| 237 |
-
g_loss = (recon_weight * recon_loss) + (perc_weight * perc_loss) + (kld_weight * kld_loss) + (
|
| 238 |
-
adversarial_weight * g_loss_adv)
|
| 239 |
-
|
| 240 |
-
g_loss.backward()
|
| 241 |
-
self.optimizer_g.step()
|
| 242 |
-
total_g_loss += g_loss.item()
|
| 243 |
-
|
| 244 |
-
avg_g_loss = total_g_loss / len(train_data) if train_data else 0.0
|
| 245 |
-
avg_d_loss = total_d_loss / len(train_data) if train_data else 0.0
|
| 246 |
-
|
| 247 |
-
print(f"Epoch {self.epoch_counter.item()} - G_Loss: {avg_g_loss:.4f}, D_Loss: {avg_d_loss:.4f}")
|
| 248 |
-
return avg_g_loss
|
| 249 |
-
|
| 250 |
-
@torch.no_grad()
|
| 251 |
-
def generate(self, text_prompts):
|
| 252 |
-
self.eval()
|
| 253 |
-
num_images = len(text_prompts)
|
| 254 |
-
z = torch.randn(num_images, self.latent_dim, device=self.device)
|
| 255 |
-
text_embeddings = self.text_encoder(text_prompts)
|
| 256 |
-
generated_images = self.cvae.decode(z, text_embeddings)
|
| 257 |
-
generated_images = (generated_images + 1) / 2
|
| 258 |
-
return [T.ToPILImage()(img.cpu()) for img in generated_images]
|
| 259 |
-
|
| 260 |
-
@torch.no_grad()
|
| 261 |
-
def forward(self, images, **kwargs):
|
| 262 |
-
prompts_to_use = kwargs.get('prompts')
|
| 263 |
-
if not prompts_to_use:
|
| 264 |
-
batch_size = images.size(0)
|
| 265 |
-
default_prompts = ["a photo of a car"]
|
| 266 |
-
prompts_to_use = [default_prompts[i % len(default_prompts)] for i in range(batch_size)]
|
| 267 |
-
|
| 268 |
-
return self.generate(prompts_to_use), prompts_to_use
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/ConvNeXt-dda5bf19-9ac1-460b-9bfd-735eec2f4904.py
DELETED
|
@@ -1,172 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
from functools import partial
|
| 3 |
-
from typing import Callable, List, Optional, Sequence
|
| 4 |
-
|
| 5 |
-
import torch
|
| 6 |
-
from torch import nn, Tensor
|
| 7 |
-
from torch.nn import functional as F
|
| 8 |
-
from torchvision.ops.misc import Conv2dNormActivation, Permute
|
| 9 |
-
from torchvision.ops.stochastic_depth import StochasticDepth
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
class LayerNorm2d(nn.LayerNorm):
|
| 13 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 14 |
-
x = x.permute(0, 2, 3, 1)
|
| 15 |
-
x = F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
|
| 16 |
-
x = x.permute(0, 3, 1, 2)
|
| 17 |
-
return x
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
class CNBlock(nn.Module):
|
| 21 |
-
def __init__(
|
| 22 |
-
self,
|
| 23 |
-
dim,
|
| 24 |
-
layer_scale: float,
|
| 25 |
-
stochastic_depth_prob: float,
|
| 26 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 27 |
-
) -> None:
|
| 28 |
-
super().__init__()
|
| 29 |
-
if norm_layer is None:
|
| 30 |
-
norm_layer = partial(nn.LayerNorm, eps=1e-6)
|
| 31 |
-
|
| 32 |
-
self.block = nn.Sequential(
|
| 33 |
-
nn.Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim, bias=True),
|
| 34 |
-
Permute([0, 2, 3, 1]),
|
| 35 |
-
norm_layer(dim),
|
| 36 |
-
nn.Linear(in_features=dim, out_features=4 * dim, bias=True),
|
| 37 |
-
nn.GELU(),
|
| 38 |
-
nn.Linear(in_features=4 * dim, out_features=dim, bias=True),
|
| 39 |
-
Permute([0, 3, 1, 2]),
|
| 40 |
-
)
|
| 41 |
-
self.layer_scale = nn.Parameter(torch.ones(dim, 1, 1) * layer_scale)
|
| 42 |
-
self.stochastic_depth = StochasticDepth(stochastic_depth_prob, "row")
|
| 43 |
-
|
| 44 |
-
def forward(self, input: Tensor) -> Tensor:
|
| 45 |
-
result = self.layer_scale * self.block(input)
|
| 46 |
-
result = self.stochastic_depth(result)
|
| 47 |
-
result += input
|
| 48 |
-
return result
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
class CNBlockConfig:
|
| 52 |
-
def __init__(
|
| 53 |
-
self,
|
| 54 |
-
input_channels: int,
|
| 55 |
-
out_channels: Optional[int],
|
| 56 |
-
num_layers: int,
|
| 57 |
-
) -> None:
|
| 58 |
-
self.input_channels = input_channels
|
| 59 |
-
self.out_channels = out_channels
|
| 60 |
-
self.num_layers = num_layers
|
| 61 |
-
|
| 62 |
-
def __repr__(self) -> str:
|
| 63 |
-
s = self.__class__.__name__ + "("
|
| 64 |
-
s += "input_channels={input_channels}"
|
| 65 |
-
s += ", out_channels={out_channels}"
|
| 66 |
-
s += ", num_layers={num_layers}"
|
| 67 |
-
s += ")"
|
| 68 |
-
return s.format(**self.__dict__)
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
def supported_hyperparameters():
|
| 72 |
-
return {'lr', 'momentum', 'stochastic_depth_prob', 'norm_eps', 'norm_std'}
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
class Net(nn.Module):
|
| 76 |
-
|
| 77 |
-
def train_setup(self, prm):
|
| 78 |
-
self.to(self.device)
|
| 79 |
-
self.criteria = (nn.CrossEntropyLoss().to(self.device),)
|
| 80 |
-
self.optimizer = torch.optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 81 |
-
|
| 82 |
-
def learn(self, train_data):
|
| 83 |
-
for inputs, labels in train_data:
|
| 84 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 85 |
-
self.optimizer.zero_grad()
|
| 86 |
-
outputs = self(inputs)
|
| 87 |
-
loss = self.criteria[0](outputs, labels)
|
| 88 |
-
loss.backward()
|
| 89 |
-
self.optimizer.step()
|
| 90 |
-
|
| 91 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 92 |
-
super().__init__()
|
| 93 |
-
self.device = device
|
| 94 |
-
num_classes: int = out_shape[0]
|
| 95 |
-
stochastic_depth_prob: float = prm['stochastic_depth_prob']
|
| 96 |
-
layer_scale: float = 1e-6
|
| 97 |
-
block_setting = None
|
| 98 |
-
block: Optional[Callable[..., nn.Module]] = None
|
| 99 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None
|
| 100 |
-
if block_setting is None:
|
| 101 |
-
block_setting = [
|
| 102 |
-
CNBlockConfig(96, 192, 4), # Changed num_layers from 3 to 4
|
| 103 |
-
CNBlockConfig(192, 384, 3),
|
| 104 |
-
CNBlockConfig(384, 768, 27),
|
| 105 |
-
CNBlockConfig(768, None, 3),
|
| 106 |
-
]
|
| 107 |
-
if not block_setting:
|
| 108 |
-
raise ValueError("The block_setting should not be empty")
|
| 109 |
-
elif not (isinstance(block_setting, Sequence) and all([isinstance(s, CNBlockConfig) for s in block_setting])):
|
| 110 |
-
raise TypeError("The block_setting should be List[CNBlockConfig]")
|
| 111 |
-
|
| 112 |
-
if block is None:
|
| 113 |
-
block = CNBlock
|
| 114 |
-
if norm_layer is None:
|
| 115 |
-
norm_layer = partial(LayerNorm2d, eps=prm['norm_eps'])
|
| 116 |
-
layers: List[nn.Module] = []
|
| 117 |
-
firstconv_output_channels = block_setting[0].input_channels
|
| 118 |
-
layers.append(
|
| 119 |
-
Conv2dNormActivation(
|
| 120 |
-
in_shape[1],
|
| 121 |
-
firstconv_output_channels,
|
| 122 |
-
kernel_size=4, # Changed kernel_size from 4 to 5
|
| 123 |
-
stride=4,
|
| 124 |
-
padding=0,
|
| 125 |
-
norm_layer=norm_layer,
|
| 126 |
-
activation_layer=None,
|
| 127 |
-
bias=True,
|
| 128 |
-
)
|
| 129 |
-
)
|
| 130 |
-
|
| 131 |
-
total_stage_blocks = sum(cnf.num_layers for cnf in block_setting)
|
| 132 |
-
stage_block_id = 0
|
| 133 |
-
for cnf in block_setting:
|
| 134 |
-
stage: List[nn.Module] = []
|
| 135 |
-
for _ in range(cnf.num_layers):
|
| 136 |
-
sd_prob = stochastic_depth_prob * stage_block_id / (total_stage_blocks - 1.0)
|
| 137 |
-
stage.append(block(cnf.input_channels, layer_scale, sd_prob))
|
| 138 |
-
stage_block_id += 1
|
| 139 |
-
layers.append(nn.Sequential(*stage))
|
| 140 |
-
if cnf.out_channels is not None:
|
| 141 |
-
layers.append(
|
| 142 |
-
nn.Sequential(
|
| 143 |
-
norm_layer(cnf.input_channels),
|
| 144 |
-
nn.Conv2d(cnf.input_channels, cnf.out_channels, kernel_size=2, stride=2),
|
| 145 |
-
)
|
| 146 |
-
)
|
| 147 |
-
|
| 148 |
-
self.features = nn.Sequential(*layers)
|
| 149 |
-
self.avgpool = nn.AdaptiveAvgPool2d(1)
|
| 150 |
-
|
| 151 |
-
lastblock = block_setting[-1]
|
| 152 |
-
lastconv_output_channels = (
|
| 153 |
-
lastblock.out_channels if lastblock.out_channels is not None else lastblock.input_channels
|
| 154 |
-
)
|
| 155 |
-
self.classifier = nn.Sequential(
|
| 156 |
-
norm_layer(lastconv_output_channels), nn.Flatten(1), nn.Linear(lastconv_output_channels, num_classes)
|
| 157 |
-
)
|
| 158 |
-
|
| 159 |
-
for m in self.modules():
|
| 160 |
-
if isinstance(m, (nn.Conv2d, nn.Linear)):
|
| 161 |
-
nn.init.trunc_normal_(m.weight, std=prm['norm_std'])
|
| 162 |
-
if m.bias is not None:
|
| 163 |
-
nn.init.zeros_(m.bias)
|
| 164 |
-
|
| 165 |
-
def _forward_impl(self, x: Tensor) -> Tensor:
|
| 166 |
-
x = self.features(x)
|
| 167 |
-
x = self.avgpool(x)
|
| 168 |
-
x = self.classifier(x)
|
| 169 |
-
return x
|
| 170 |
-
|
| 171 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 172 |
-
return self._forward_impl(x)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN107.py
DELETED
|
@@ -1,92 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.optim as optim
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
# Define DPNBlock with Group Convolutions
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels, stride=1):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(
|
| 15 |
-
in_channels, out_channels, kernel_size=3, stride=stride, padding=1, groups=4, bias=False
|
| 16 |
-
)
|
| 17 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 18 |
-
self.relu = nn.ReLU(inplace=True)
|
| 19 |
-
self.conv2 = nn.Conv2d(
|
| 20 |
-
out_channels, out_channels, kernel_size=3, padding=1, groups=4, bias=False
|
| 21 |
-
)
|
| 22 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 23 |
-
|
| 24 |
-
def forward(self, x):
|
| 25 |
-
residual = x
|
| 26 |
-
out = self.conv1(x)
|
| 27 |
-
out = self.bn1(out)
|
| 28 |
-
out = self.relu(out)
|
| 29 |
-
out = self.conv2(out)
|
| 30 |
-
out = self.bn2(out)
|
| 31 |
-
out += residual # Residual connection
|
| 32 |
-
return self.relu(out)
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
# Memory-Optimized DPN107
|
| 36 |
-
class DPN107(nn.Module):
|
| 37 |
-
def __init__(self, in_channels, num_classes, num_blocks, growth_rate):
|
| 38 |
-
super(DPN107, self).__init__()
|
| 39 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1, bias=False)
|
| 40 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 41 |
-
self.relu = nn.ReLU(inplace=True)
|
| 42 |
-
|
| 43 |
-
self.blocks = nn.Sequential(
|
| 44 |
-
*[DPNBlock(growth_rate, growth_rate) for _ in range(num_blocks)]
|
| 45 |
-
)
|
| 46 |
-
|
| 47 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 48 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 49 |
-
|
| 50 |
-
def forward(self, x):
|
| 51 |
-
x = self.conv1(x)
|
| 52 |
-
x = self.bn1(x)
|
| 53 |
-
x = self.relu(x)
|
| 54 |
-
x = self.blocks(x)
|
| 55 |
-
x = self.avgpool(x)
|
| 56 |
-
x = torch.flatten(x, 1)
|
| 57 |
-
x = self.fc(x)
|
| 58 |
-
return x
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
class Net(nn.Module):
|
| 62 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 63 |
-
super(Net, self).__init__()
|
| 64 |
-
self.device = device
|
| 65 |
-
model_class = DPN107
|
| 66 |
-
self.channel_number = in_shape[1]
|
| 67 |
-
self.image_size = in_shape[2]
|
| 68 |
-
self.class_number = out_shape[0]
|
| 69 |
-
|
| 70 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=3, growth_rate=32)
|
| 71 |
-
self.learning_rate = prm['lr']
|
| 72 |
-
self.momentum = prm['momentum']
|
| 73 |
-
|
| 74 |
-
def forward(self, x):
|
| 75 |
-
return self.model(x)
|
| 76 |
-
|
| 77 |
-
def train_setup(self, prm):
|
| 78 |
-
self.to(self.device)
|
| 79 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 80 |
-
self.optimizer = optim.SGD(self.parameters(), lr=self.learning_rate, momentum=self.momentum)
|
| 81 |
-
|
| 82 |
-
def learn(self, train_data):
|
| 83 |
-
self.train()
|
| 84 |
-
for inputs, labels in train_data:
|
| 85 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 86 |
-
inputs = inputs.float()
|
| 87 |
-
self.optimizer.zero_grad()
|
| 88 |
-
outputs = self(inputs)
|
| 89 |
-
loss = self.criteria(outputs, labels)
|
| 90 |
-
loss.backward()
|
| 91 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 92 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN131-8e6e495b-85cb-4a71-8b91-6d89372e0a0c.py
DELETED
|
@@ -1,86 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels, stride=1):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1)
|
| 15 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 16 |
-
self.relu = nn.ReLU(inplace=True)
|
| 17 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
| 18 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
residual = x
|
| 22 |
-
out = self.conv1(x)
|
| 23 |
-
out = self.bn1(out)
|
| 24 |
-
out = self.relu(out)
|
| 25 |
-
out = self.conv2(out)
|
| 26 |
-
out = self.bn2(out)
|
| 27 |
-
out += residual
|
| 28 |
-
return self.relu(out)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class DPN131(nn.Module):
|
| 32 |
-
def __init__(self, in_channels=3, num_classes=10, num_blocks=4, growth_rate=35):
|
| 33 |
-
super(DPN131, self).__init__()
|
| 34 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1)
|
| 35 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 36 |
-
self.relu = nn.ReLU(inplace=True)
|
| 37 |
-
|
| 38 |
-
self.blocks = nn.ModuleList()
|
| 39 |
-
for _ in range(num_blocks):
|
| 40 |
-
self.blocks.append(DPNBlock(growth_rate, growth_rate))
|
| 41 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 42 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 43 |
-
|
| 44 |
-
def forward(self, x):
|
| 45 |
-
x = self.conv1(x)
|
| 46 |
-
x = self.bn1(x)
|
| 47 |
-
x = self.relu(x)
|
| 48 |
-
for block in self.blocks:
|
| 49 |
-
x = block(x)
|
| 50 |
-
x = self.avgpool(x)
|
| 51 |
-
x = torch.flatten(x, 1)
|
| 52 |
-
x = self.fc(x)
|
| 53 |
-
return x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 58 |
-
super(Net, self).__init__()
|
| 59 |
-
self.device = device
|
| 60 |
-
model_class = DPN131
|
| 61 |
-
self.channel_number = in_shape[1]
|
| 62 |
-
self.image_size = in_shape[2]
|
| 63 |
-
self.class_number = out_shape[0]
|
| 64 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=3, growth_rate=32)
|
| 65 |
-
|
| 66 |
-
self.learning_rate = prm['lr']
|
| 67 |
-
self.momentum = prm['momentum']
|
| 68 |
-
|
| 69 |
-
def forward(self, x):
|
| 70 |
-
return self.model(x)
|
| 71 |
-
|
| 72 |
-
def train_setup(self, prm):
|
| 73 |
-
self.to(self.device)
|
| 74 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 75 |
-
self.optimizer = optim.SGD(self.parameters(), lr=self.learning_rate, momentum=self.momentum)
|
| 76 |
-
|
| 77 |
-
def learn(self, train_data):
|
| 78 |
-
self.train()
|
| 79 |
-
for inputs, labels in train_data:
|
| 80 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 81 |
-
self.optimizer.zero_grad()
|
| 82 |
-
outputs = self(inputs)
|
| 83 |
-
loss = self.criteria(outputs, labels)
|
| 84 |
-
loss.backward()
|
| 85 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 86 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN131-c53a40b8-b874-4c8b-999b-0944a1173a46.py
DELETED
|
@@ -1,86 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels, stride=1):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1)
|
| 15 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 16 |
-
self.relu = nn.ReLU(inplace=True)
|
| 17 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
| 18 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
residual = x
|
| 22 |
-
out = self.conv1(x)
|
| 23 |
-
out = self.bn1(out)
|
| 24 |
-
out = self.relu(out)
|
| 25 |
-
out = self.conv2(out)
|
| 26 |
-
out = self.bn2(out)
|
| 27 |
-
out += residual
|
| 28 |
-
return self.relu(out)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class DPN131(nn.Module):
|
| 32 |
-
def __init__(self, in_channels=3, num_classes=10, num_blocks=4, growth_rate=16):
|
| 33 |
-
super(DPN131, self).__init__()
|
| 34 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1)
|
| 35 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 36 |
-
self.relu = nn.ReLU(inplace=True)
|
| 37 |
-
|
| 38 |
-
self.blocks = nn.ModuleList()
|
| 39 |
-
for _ in range(num_blocks):
|
| 40 |
-
self.blocks.append(DPNBlock(growth_rate, growth_rate))
|
| 41 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 42 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 43 |
-
|
| 44 |
-
def forward(self, x):
|
| 45 |
-
x = self.conv1(x)
|
| 46 |
-
x = self.bn1(x)
|
| 47 |
-
x = self.relu(x)
|
| 48 |
-
for block in self.blocks:
|
| 49 |
-
x = block(x)
|
| 50 |
-
x = self.avgpool(x)
|
| 51 |
-
x = torch.flatten(x, 1)
|
| 52 |
-
x = self.fc(x)
|
| 53 |
-
return x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 58 |
-
super(Net, self).__init__()
|
| 59 |
-
self.device = device
|
| 60 |
-
model_class = DPN131
|
| 61 |
-
self.channel_number = in_shape[1]
|
| 62 |
-
self.image_size = in_shape[2]
|
| 63 |
-
self.class_number = out_shape[0]
|
| 64 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=4, growth_rate=16)
|
| 65 |
-
|
| 66 |
-
self.learning_rate = prm['lr']
|
| 67 |
-
self.momentum = prm['momentum']
|
| 68 |
-
|
| 69 |
-
def forward(self, x):
|
| 70 |
-
return self.model(x)
|
| 71 |
-
|
| 72 |
-
def train_setup(self, prm):
|
| 73 |
-
self.to(self.device)
|
| 74 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 75 |
-
self.optimizer = optim.SGD(self.parameters(), lr=self.learning_rate, momentum=self.momentum)
|
| 76 |
-
|
| 77 |
-
def learn(self, train_data):
|
| 78 |
-
self.train()
|
| 79 |
-
for inputs, labels in train_data:
|
| 80 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 81 |
-
self.optimizer.zero_grad()
|
| 82 |
-
outputs = self(inputs)
|
| 83 |
-
loss = self.criteria(outputs, labels)
|
| 84 |
-
loss.backward()
|
| 85 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 86 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN131-e8980802-6b89-4170-8608-327297706df0.py
DELETED
|
@@ -1,86 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels, stride=1):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1)
|
| 15 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 16 |
-
self.relu = nn.ReLU(inplace=True)
|
| 17 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
| 18 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
residual = x
|
| 22 |
-
out = self.conv1(x)
|
| 23 |
-
out = self.bn1(out)
|
| 24 |
-
out = self.relu(out)
|
| 25 |
-
out = self.conv2(out)
|
| 26 |
-
out = self.bn2(out)
|
| 27 |
-
out += residual
|
| 28 |
-
return self.relu(out)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class DPN131(nn.Module):
|
| 32 |
-
def __init__(self, in_channels=3, num_classes=10, num_blocks=5, growth_rate=40):
|
| 33 |
-
super(DPN131, self).__init__()
|
| 34 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1)
|
| 35 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 36 |
-
self.relu = nn.ReLU(inplace=True)
|
| 37 |
-
|
| 38 |
-
self.blocks = nn.ModuleList()
|
| 39 |
-
for _ in range(num_blocks):
|
| 40 |
-
self.blocks.append(DPNBlock(growth_rate, growth_rate))
|
| 41 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 42 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 43 |
-
|
| 44 |
-
def forward(self, x):
|
| 45 |
-
x = self.conv1(x)
|
| 46 |
-
x = self.bn1(x)
|
| 47 |
-
x = self.relu(x)
|
| 48 |
-
for block in self.blocks:
|
| 49 |
-
x = block(x)
|
| 50 |
-
x = self.avgpool(x)
|
| 51 |
-
x = torch.flatten(x, 1)
|
| 52 |
-
x = self.fc(x)
|
| 53 |
-
return x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 58 |
-
super(Net, self).__init__()
|
| 59 |
-
self.device = device
|
| 60 |
-
model_class = DPN131
|
| 61 |
-
self.channel_number = in_shape[1]
|
| 62 |
-
self.image_size = in_shape[2]
|
| 63 |
-
self.class_number = out_shape[0]
|
| 64 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=3, growth_rate=32)
|
| 65 |
-
|
| 66 |
-
self.learning_rate = prm['lr']
|
| 67 |
-
self.momentum = prm['momentum']
|
| 68 |
-
|
| 69 |
-
def forward(self, x):
|
| 70 |
-
return self.model(x)
|
| 71 |
-
|
| 72 |
-
def train_setup(self, prm):
|
| 73 |
-
self.to(self.device)
|
| 74 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 75 |
-
self.optimizer = optim.SGD(self.parameters(), lr=self.learning_rate, momentum=self.momentum)
|
| 76 |
-
|
| 77 |
-
def learn(self, train_data):
|
| 78 |
-
self.train()
|
| 79 |
-
for inputs, labels in train_data:
|
| 80 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 81 |
-
self.optimizer.zero_grad()
|
| 82 |
-
outputs = self(inputs)
|
| 83 |
-
loss = self.criteria(outputs, labels)
|
| 84 |
-
loss.backward()
|
| 85 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 86 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN131.py
DELETED
|
@@ -1,85 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.optim as optim
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class DPNBlock(nn.Module):
|
| 11 |
-
def __init__(self, in_channels, out_channels, stride=1):
|
| 12 |
-
super(DPNBlock, self).__init__()
|
| 13 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1)
|
| 14 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 15 |
-
self.relu = nn.ReLU(inplace=True)
|
| 16 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
| 17 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 18 |
-
|
| 19 |
-
def forward(self, x):
|
| 20 |
-
residual = x
|
| 21 |
-
out = self.conv1(x)
|
| 22 |
-
out = self.bn1(out)
|
| 23 |
-
out = self.relu(out)
|
| 24 |
-
out = self.conv2(out)
|
| 25 |
-
out = self.bn2(out)
|
| 26 |
-
out += residual
|
| 27 |
-
return self.relu(out)
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
class DPN131(nn.Module):
|
| 31 |
-
def __init__(self, in_channels=3, num_classes=10, num_blocks=3, growth_rate=32):
|
| 32 |
-
super(DPN131, self).__init__()
|
| 33 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, padding=1)
|
| 34 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 35 |
-
self.relu = nn.ReLU(inplace=True)
|
| 36 |
-
|
| 37 |
-
self.blocks = nn.ModuleList()
|
| 38 |
-
for _ in range(num_blocks):
|
| 39 |
-
self.blocks.append(DPNBlock(growth_rate, growth_rate))
|
| 40 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 41 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 42 |
-
|
| 43 |
-
def forward(self, x):
|
| 44 |
-
x = self.conv1(x)
|
| 45 |
-
x = self.bn1(x)
|
| 46 |
-
x = self.relu(x)
|
| 47 |
-
for block in self.blocks:
|
| 48 |
-
x = block(x)
|
| 49 |
-
x = self.avgpool(x)
|
| 50 |
-
x = torch.flatten(x, 1)
|
| 51 |
-
x = self.fc(x)
|
| 52 |
-
return x
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
class Net(nn.Module):
|
| 56 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 57 |
-
super(Net, self).__init__()
|
| 58 |
-
self.device = device
|
| 59 |
-
model_class = DPN131
|
| 60 |
-
self.channel_number = in_shape[1]
|
| 61 |
-
self.image_size = in_shape[2]
|
| 62 |
-
self.class_number = out_shape[0]
|
| 63 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=3, growth_rate=32)
|
| 64 |
-
|
| 65 |
-
self.learning_rate = prm['lr']
|
| 66 |
-
self.momentum = prm['momentum']
|
| 67 |
-
|
| 68 |
-
def forward(self, x):
|
| 69 |
-
return self.model(x)
|
| 70 |
-
|
| 71 |
-
def train_setup(self, prm):
|
| 72 |
-
self.to(self.device)
|
| 73 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 74 |
-
self.optimizer = optim.SGD(self.parameters(), lr=self.learning_rate, momentum=self.momentum)
|
| 75 |
-
|
| 76 |
-
def learn(self, train_data):
|
| 77 |
-
self.train()
|
| 78 |
-
for inputs, labels in train_data:
|
| 79 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 80 |
-
self.optimizer.zero_grad()
|
| 81 |
-
outputs = self(inputs)
|
| 82 |
-
loss = self.criteria(outputs, labels)
|
| 83 |
-
loss.backward()
|
| 84 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 85 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN68-9693aa0b-80bf-4393-a9e0-dd985a5ab128.py
DELETED
|
@@ -1,83 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=5, stride=1, padding=2)
|
| 15 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 16 |
-
self.relu = nn.ReLU(inplace=True)
|
| 17 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
| 18 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
residual = x
|
| 22 |
-
out = self.conv1(x)
|
| 23 |
-
out = self.bn1(out)
|
| 24 |
-
out = self.relu(out)
|
| 25 |
-
out = self.conv2(out)
|
| 26 |
-
out = self.bn2(out)
|
| 27 |
-
out = out + residual
|
| 28 |
-
return self.relu(out)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class DPN68(nn.Module):
|
| 32 |
-
def __init__(self, in_channels, num_classes, num_blocks, growth_rate):
|
| 33 |
-
super(DPN68, self).__init__()
|
| 34 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, stride=1, padding=1)
|
| 35 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 36 |
-
self.relu = nn.ReLU(inplace=True)
|
| 37 |
-
|
| 38 |
-
self.blocks = nn.Sequential(
|
| 39 |
-
*[DPNBlock(growth_rate, growth_rate) for _ in range(num_blocks)]
|
| 40 |
-
)
|
| 41 |
-
|
| 42 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 43 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 44 |
-
|
| 45 |
-
def forward(self, x):
|
| 46 |
-
x = self.conv1(x)
|
| 47 |
-
x = self.bn1(x)
|
| 48 |
-
x = self.relu(x)
|
| 49 |
-
x = self.blocks(x)
|
| 50 |
-
x = self.avgpool(x)
|
| 51 |
-
x = torch.flatten(x, 1)
|
| 52 |
-
x = self.fc(x)
|
| 53 |
-
return x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 58 |
-
super(Net, self).__init__()
|
| 59 |
-
self.device = device
|
| 60 |
-
model_class = DPN68
|
| 61 |
-
self.channel_number = in_shape[1]
|
| 62 |
-
self.image_size = in_shape[2]
|
| 63 |
-
self.class_number = out_shape[0]
|
| 64 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=5, growth_rate=16)
|
| 65 |
-
|
| 66 |
-
def forward(self, x):
|
| 67 |
-
return self.model(x)
|
| 68 |
-
|
| 69 |
-
def train_setup(self, prm):
|
| 70 |
-
self.to(self.device)
|
| 71 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 72 |
-
self.optimizer = optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 73 |
-
|
| 74 |
-
def learn(self, train_data):
|
| 75 |
-
self.train()
|
| 76 |
-
for inputs, labels in train_data:
|
| 77 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 78 |
-
self.optimizer.zero_grad()
|
| 79 |
-
outputs = self(inputs)
|
| 80 |
-
loss = self.criteria(outputs, labels)
|
| 81 |
-
loss.backward()
|
| 82 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 83 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DPN68-c9cdb196-7596-4368-974a-56edf8b10381.py
DELETED
|
@@ -1,83 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DPNBlock(nn.Module):
|
| 12 |
-
def __init__(self, in_channels, out_channels):
|
| 13 |
-
super(DPNBlock, self).__init__()
|
| 14 |
-
self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=5, stride=1, padding=2)
|
| 15 |
-
self.bn1 = nn.BatchNorm2d(out_channels)
|
| 16 |
-
self.relu = nn.ReLU(inplace=True)
|
| 17 |
-
self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
| 18 |
-
self.bn2 = nn.BatchNorm2d(out_channels)
|
| 19 |
-
|
| 20 |
-
def forward(self, x):
|
| 21 |
-
residual = x
|
| 22 |
-
out = self.conv1(x)
|
| 23 |
-
out = self.bn1(out)
|
| 24 |
-
out = self.relu(out)
|
| 25 |
-
out = self.conv2(out)
|
| 26 |
-
out = self.bn2(out)
|
| 27 |
-
out = out + residual
|
| 28 |
-
return self.relu(out)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class DPN68(nn.Module):
|
| 32 |
-
def __init__(self, in_channels, num_classes, num_blocks, growth_rate):
|
| 33 |
-
super(DPN68, self).__init__()
|
| 34 |
-
self.conv1 = nn.Conv2d(in_channels, growth_rate, kernel_size=3, stride=1, padding=1)
|
| 35 |
-
self.bn1 = nn.BatchNorm2d(growth_rate)
|
| 36 |
-
self.relu = nn.ReLU(inplace=True)
|
| 37 |
-
|
| 38 |
-
self.blocks = nn.Sequential(
|
| 39 |
-
*[DPNBlock(growth_rate, growth_rate) for _ in range(num_blocks)]
|
| 40 |
-
)
|
| 41 |
-
|
| 42 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 43 |
-
self.fc = nn.Linear(growth_rate, num_classes)
|
| 44 |
-
|
| 45 |
-
def forward(self, x):
|
| 46 |
-
x = self.conv1(x)
|
| 47 |
-
x = self.bn1(x)
|
| 48 |
-
x = self.relu(x)
|
| 49 |
-
x = self.blocks(x)
|
| 50 |
-
x = self.avgpool(x)
|
| 51 |
-
x = torch.flatten(x, 1)
|
| 52 |
-
x = self.fc(x)
|
| 53 |
-
return x
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
class Net(nn.Module):
|
| 57 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 58 |
-
super(Net, self).__init__()
|
| 59 |
-
self.device = device
|
| 60 |
-
model_class = DPN68
|
| 61 |
-
self.channel_number = in_shape[1]
|
| 62 |
-
self.image_size = in_shape[2]
|
| 63 |
-
self.class_number = out_shape[0]
|
| 64 |
-
self.model = model_class(self.channel_number, self.class_number, num_blocks=3, growth_rate=32)
|
| 65 |
-
|
| 66 |
-
def forward(self, x):
|
| 67 |
-
return self.model(x)
|
| 68 |
-
|
| 69 |
-
def train_setup(self, prm):
|
| 70 |
-
self.to(self.device)
|
| 71 |
-
self.criteria = nn.CrossEntropyLoss().to(self.device)
|
| 72 |
-
self.optimizer = optim.SGD(self.parameters(), lr=prm['lr'], momentum=prm['momentum'])
|
| 73 |
-
|
| 74 |
-
def learn(self, train_data):
|
| 75 |
-
self.train()
|
| 76 |
-
for inputs, labels in train_data:
|
| 77 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 78 |
-
self.optimizer.zero_grad()
|
| 79 |
-
outputs = self(inputs)
|
| 80 |
-
loss = self.criteria(outputs, labels)
|
| 81 |
-
loss.backward()
|
| 82 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 83 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DarkNet-11e8caec-5e73-461e-a101-3aa39dfec644.py
DELETED
|
@@ -1,96 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum', 'dropout'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DarkNetUnit(nn.Module):
|
| 12 |
-
def __init__(self, in_channels: int, out_channels: int, pointwise: bool, alpha: float):
|
| 13 |
-
super(DarkNetUnit, self).__init__()
|
| 14 |
-
self.activation = nn.LeakyReLU(negative_slope=alpha, inplace=True)
|
| 15 |
-
if pointwise:
|
| 16 |
-
self.conv = nn.Sequential(
|
| 17 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, bias=False),
|
| 18 |
-
nn.BatchNorm2d(out_channels),
|
| 19 |
-
self.activation
|
| 20 |
-
)
|
| 21 |
-
else:
|
| 22 |
-
self.conv = nn.Sequential(
|
| 23 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False),
|
| 24 |
-
nn.BatchNorm2d(out_channels),
|
| 25 |
-
self.activation
|
| 26 |
-
)
|
| 27 |
-
|
| 28 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 29 |
-
return self.conv(x)
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
class Net(nn.Module):
|
| 33 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 34 |
-
super(Net, self).__init__()
|
| 35 |
-
self.device = device
|
| 36 |
-
channels: list = [[32, 32, 32], [192, 192, 192], [128, 128, 128], [256, 256, 256]]
|
| 37 |
-
odd_pointwise: bool = True
|
| 38 |
-
alpha: float = 0.1
|
| 39 |
-
in_channels = in_shape[1]
|
| 40 |
-
image_size = in_shape[2]
|
| 41 |
-
num_classes = out_shape[0]
|
| 42 |
-
|
| 43 |
-
if channels is None:
|
| 44 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 45 |
-
|
| 46 |
-
self.features = nn.Sequential()
|
| 47 |
-
for i, channels_per_stage in enumerate(channels):
|
| 48 |
-
stage = nn.Sequential()
|
| 49 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 50 |
-
pointwise = (len(channels_per_stage) > 1) and not (((j + 1) % 2 == 1) ^ odd_pointwise)
|
| 51 |
-
stage.add_module(f"unit{j + 1}", DarkNetUnit(in_channels, out_channels, pointwise, alpha))
|
| 52 |
-
in_channels = out_channels
|
| 53 |
-
if i != len(channels) - 1:
|
| 54 |
-
stage.add_module(f"pool{i + 1}", nn.MaxPool2d(kernel_size=2, stride=2))
|
| 55 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 56 |
-
|
| 57 |
-
final_feature_map_size = image_size // (2 ** (len(channels) - 1))
|
| 58 |
-
|
| 59 |
-
self.output = nn.Sequential(
|
| 60 |
-
nn.Conv2d(in_channels=in_channels, out_channels=num_classes, kernel_size=1),
|
| 61 |
-
nn.LeakyReLU(negative_slope=alpha, inplace=True),
|
| 62 |
-
nn.AdaptiveAvgPool2d(output_size=(1, 1))
|
| 63 |
-
)
|
| 64 |
-
|
| 65 |
-
self._initialize_weights()
|
| 66 |
-
|
| 67 |
-
def _initialize_weights(self):
|
| 68 |
-
for module in self.modules():
|
| 69 |
-
if isinstance(module, nn.Conv2d):
|
| 70 |
-
nn.init.kaiming_uniform_(module.weight, mode='fan_in', nonlinearity='leaky_relu')
|
| 71 |
-
if module.bias is not None:
|
| 72 |
-
nn.init.constant_(module.bias, 0)
|
| 73 |
-
|
| 74 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 75 |
-
x = self.features(x)
|
| 76 |
-
x = self.output(x)
|
| 77 |
-
x = x.view(x.size(0), -1)
|
| 78 |
-
return x
|
| 79 |
-
|
| 80 |
-
def train_setup(self, prm: dict):
|
| 81 |
-
self.to(self.device)
|
| 82 |
-
learning_rate = float(prm.get("lr", 0.01))
|
| 83 |
-
momentum = float(prm.get("momentum", 0.9))
|
| 84 |
-
self.criteria = nn.CrossEntropyLoss()
|
| 85 |
-
self.optimizer = optim.SGD(self.parameters(), lr=learning_rate, momentum=momentum)
|
| 86 |
-
self.to(self.device)
|
| 87 |
-
|
| 88 |
-
def learn(self, train_data: torch.utils.data.DataLoader):
|
| 89 |
-
self.train()
|
| 90 |
-
for inputs, targets in train_data:
|
| 91 |
-
inputs, targets = inputs.to(next(self.parameters()).device), targets.to(next(self.parameters()).device)
|
| 92 |
-
self.optimizer.zero_grad()
|
| 93 |
-
outputs = self(inputs)
|
| 94 |
-
loss = self.criteria(outputs, targets)
|
| 95 |
-
loss.backward()
|
| 96 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DarkNet-51277b91-c3a1-4669-9fb5-849ea97bd1b4.py
DELETED
|
@@ -1,95 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum', 'dropout'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DarkNetUnit(nn.Module):
|
| 12 |
-
def __init__(self, in_channels: int, out_channels: int, pointwise: bool, alpha: float):
|
| 13 |
-
super(DarkNetUnit, self).__init__()
|
| 14 |
-
self.activation = nn.LeakyReLU(negative_slope=alpha, inplace=True)
|
| 15 |
-
if pointwise:
|
| 16 |
-
self.conv = nn.Sequential(
|
| 17 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, bias=False),
|
| 18 |
-
nn.BatchNorm2d(out_channels),
|
| 19 |
-
self.activation
|
| 20 |
-
)
|
| 21 |
-
else:
|
| 22 |
-
self.conv = nn.Sequential(
|
| 23 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False),
|
| 24 |
-
nn.BatchNorm2d(out_channels),
|
| 25 |
-
self.activation
|
| 26 |
-
)
|
| 27 |
-
|
| 28 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 29 |
-
return self.conv(x)
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
class Net(nn.Module):
|
| 33 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 34 |
-
super(Net, self).__init__()
|
| 35 |
-
self.device = device
|
| 36 |
-
channels: list = None
|
| 37 |
-
odd_pointwise: bool = True
|
| 38 |
-
alpha: float = 0.2 # Changed from 0.1
|
| 39 |
-
in_channels = in_shape[1]
|
| 40 |
-
image_size = 200 # Changed from in_shape[2]
|
| 41 |
-
num_classes = 1000 # Changed from out_shape[0]
|
| 42 |
-
# Changed first element of channels list
|
| 43 |
-
channels = [[32, 32, 32], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 44 |
-
|
| 45 |
-
self.features = nn.Sequential()
|
| 46 |
-
for i, channels_per_stage in enumerate(channels):
|
| 47 |
-
stage = nn.Sequential()
|
| 48 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 49 |
-
pointwise = (len(channels_per_stage) > 1) and not (((j + 1) % 2 == 1) ^ odd_pointwise)
|
| 50 |
-
stage.add_module(f"unit{j + 1}", DarkNetUnit(in_channels, out_channels, pointwise, alpha))
|
| 51 |
-
in_channels = out_channels
|
| 52 |
-
if i != len(channels) - 1:
|
| 53 |
-
stage.add_module(f"pool{i + 1}", nn.MaxPool2d(kernel_size=2, stride=2))
|
| 54 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 55 |
-
|
| 56 |
-
final_feature_map_size = image_size // (2 ** (len(channels) - 1))
|
| 57 |
-
|
| 58 |
-
self.output = nn.Sequential(
|
| 59 |
-
nn.Conv2d(in_channels=in_channels, out_channels=num_classes, kernel_size=1),
|
| 60 |
-
nn.LeakyReLU(negative_slope=alpha, inplace=True),
|
| 61 |
-
nn.AdaptiveAvgPool2d(output_size=(1, 1))
|
| 62 |
-
)
|
| 63 |
-
|
| 64 |
-
self._initialize_weights()
|
| 65 |
-
|
| 66 |
-
def _initialize_weights(self):
|
| 67 |
-
for module in self.modules():
|
| 68 |
-
if isinstance(module, nn.Conv2d):
|
| 69 |
-
nn.init.kaiming_uniform_(module.weight, mode='fan_in', nonlinearity='leaky_relu')
|
| 70 |
-
if module.bias is not None:
|
| 71 |
-
nn.init.constant_(module.bias, 0)
|
| 72 |
-
|
| 73 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 74 |
-
x = self.features(x)
|
| 75 |
-
x = self.output(x)
|
| 76 |
-
x = x.view(x.size(0), -1)
|
| 77 |
-
return x
|
| 78 |
-
|
| 79 |
-
def train_setup(self, prm: dict):
|
| 80 |
-
self.to(self.device)
|
| 81 |
-
learning_rate = float(prm.get("lr", 0.01))
|
| 82 |
-
momentum = float(prm.get("momentum", 0.9)) # Changed from 0.9
|
| 83 |
-
self.criteria = nn.CrossEntropyLoss()
|
| 84 |
-
self.optimizer = optim.SGD(self.parameters(), lr=learning_rate, momentum=momentum)
|
| 85 |
-
self.to(self.device)
|
| 86 |
-
|
| 87 |
-
def learn(self, train_data: torch.utils.data.DataLoader):
|
| 88 |
-
self.train()
|
| 89 |
-
for inputs, targets in train_data:
|
| 90 |
-
inputs, targets = inputs.to(next(self.parameters()).device), targets.to(next(self.parameters()).device)
|
| 91 |
-
self.optimizer.zero_grad()
|
| 92 |
-
outputs = self(inputs)
|
| 93 |
-
loss = self.criteria(outputs, targets)
|
| 94 |
-
loss.backward()
|
| 95 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DarkNet-d434ba1c-25ea-4160-a41d-4c477dba7bc0.py
DELETED
|
@@ -1,96 +0,0 @@
|
|
| 1 |
-
|
| 2 |
-
import torch
|
| 3 |
-
import torch.nn as nn
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
def supported_hyperparameters():
|
| 8 |
-
return {'lr', 'momentum', 'dropout'}
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
class DarkNetUnit(nn.Module):
|
| 12 |
-
def __init__(self, in_channels: int, out_channels: int, pointwise: bool, alpha: float):
|
| 13 |
-
super(DarkNetUnit, self).__init__()
|
| 14 |
-
self.activation = nn.LeakyReLU(negative_slope=alpha, inplace=True)
|
| 15 |
-
if pointwise:
|
| 16 |
-
self.conv = nn.Sequential(
|
| 17 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, bias=False),
|
| 18 |
-
nn.BatchNorm2d(out_channels),
|
| 19 |
-
self.activation
|
| 20 |
-
)
|
| 21 |
-
else:
|
| 22 |
-
self.conv = nn.Sequential(
|
| 23 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False),
|
| 24 |
-
nn.BatchNorm2d(out_channels),
|
| 25 |
-
self.activation
|
| 26 |
-
)
|
| 27 |
-
|
| 28 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 29 |
-
return self.conv(x)
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
class Net(nn.Module):
|
| 33 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 34 |
-
super(Net, self).__init__()
|
| 35 |
-
self.device = device
|
| 36 |
-
channels: list = None
|
| 37 |
-
odd_pointwise: bool = True
|
| 38 |
-
alpha: float = 0.1
|
| 39 |
-
in_channels = in_shape[1]
|
| 40 |
-
image_size = in_shape[2]
|
| 41 |
-
num_classes = out_shape[0]
|
| 42 |
-
|
| 43 |
-
if channels is None:
|
| 44 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256]] # Changed from 4 to 3 sublists
|
| 45 |
-
|
| 46 |
-
self.features = nn.Sequential()
|
| 47 |
-
for i, channels_per_stage in enumerate(channels):
|
| 48 |
-
stage = nn.Sequential()
|
| 49 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 50 |
-
pointwise = (len(channels_per_stage) > 1) and not (((j + 1) % 2 == 1) ^ odd_pointwise)
|
| 51 |
-
stage.add_module(f"unit{j + 1}", DarkNetUnit(in_channels, out_channels, pointwise, alpha))
|
| 52 |
-
in_channels = out_channels
|
| 53 |
-
if i != len(channels) - 1:
|
| 54 |
-
stage.add_module(f"pool{i + 1}", nn.MaxPool2d(kernel_size=2, stride=2)) # Changed kernel size from 2 to 3
|
| 55 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 56 |
-
|
| 57 |
-
final_feature_map_size = image_size // (2 ** (len(channels) - 1))
|
| 58 |
-
|
| 59 |
-
self.output = nn.Sequential(
|
| 60 |
-
nn.Conv2d(in_channels=in_channels, out_channels=num_classes, kernel_size=1),
|
| 61 |
-
nn.LeakyReLU(negative_slope=alpha, inplace=True),
|
| 62 |
-
nn.AdaptiveAvgPool2d(output_size=(1, 1))
|
| 63 |
-
)
|
| 64 |
-
|
| 65 |
-
self._initialize_weights()
|
| 66 |
-
|
| 67 |
-
def _initialize_weights(self):
|
| 68 |
-
for module in self.modules():
|
| 69 |
-
if isinstance(module, nn.Conv2d):
|
| 70 |
-
nn.init.kaiming_uniform_(module.weight, mode='fan_in', nonlinearity='leaky_relu')
|
| 71 |
-
if module.bias is not None:
|
| 72 |
-
nn.init.constant_(module.bias, 0)
|
| 73 |
-
|
| 74 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 75 |
-
x = self.features(x)
|
| 76 |
-
x = self.output(x)
|
| 77 |
-
x = x.view(x.size(0), -1)
|
| 78 |
-
return x
|
| 79 |
-
|
| 80 |
-
def train_setup(self, prm: dict):
|
| 81 |
-
self.to(self.device)
|
| 82 |
-
learning_rate = float(prm.get("lr", 0.01))
|
| 83 |
-
momentum = float(prm.get("momentum", 0.9))
|
| 84 |
-
self.criteria = nn.CrossEntropyLoss()
|
| 85 |
-
self.optimizer = optim.SGD(self.parameters(), lr=learning_rate, momentum=momentum)
|
| 86 |
-
self.to(self.device)
|
| 87 |
-
|
| 88 |
-
def learn(self, train_data: torch.utils.data.DataLoader):
|
| 89 |
-
self.train()
|
| 90 |
-
for inputs, targets in train_data:
|
| 91 |
-
inputs, targets = inputs.to(next(self.parameters()).device), targets.to(next(self.parameters()).device)
|
| 92 |
-
self.optimizer.zero_grad()
|
| 93 |
-
outputs = self(inputs)
|
| 94 |
-
loss = self.criteria(outputs, targets)
|
| 95 |
-
loss.backward()
|
| 96 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DarkNet.py
DELETED
|
@@ -1,95 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.optim as optim
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr', 'momentum', 'dropout'}
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
class DarkNetUnit(nn.Module):
|
| 11 |
-
def __init__(self, in_channels: int, out_channels: int, pointwise: bool, alpha: float):
|
| 12 |
-
super(DarkNetUnit, self).__init__()
|
| 13 |
-
self.activation = nn.LeakyReLU(negative_slope=alpha, inplace=True)
|
| 14 |
-
if pointwise:
|
| 15 |
-
self.conv = nn.Sequential(
|
| 16 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, bias=False),
|
| 17 |
-
nn.BatchNorm2d(out_channels),
|
| 18 |
-
self.activation
|
| 19 |
-
)
|
| 20 |
-
else:
|
| 21 |
-
self.conv = nn.Sequential(
|
| 22 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False),
|
| 23 |
-
nn.BatchNorm2d(out_channels),
|
| 24 |
-
self.activation
|
| 25 |
-
)
|
| 26 |
-
|
| 27 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 28 |
-
return self.conv(x)
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
class Net(nn.Module):
|
| 32 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 33 |
-
super(Net, self).__init__()
|
| 34 |
-
self.device = device
|
| 35 |
-
channels: list = None
|
| 36 |
-
odd_pointwise: bool = True
|
| 37 |
-
alpha: float = 0.1
|
| 38 |
-
in_channels = in_shape[1]
|
| 39 |
-
image_size = in_shape[2]
|
| 40 |
-
num_classes = out_shape[0]
|
| 41 |
-
|
| 42 |
-
if channels is None:
|
| 43 |
-
channels = [[64, 64, 64], [128, 128, 128], [256, 256, 256], [512, 512, 512]]
|
| 44 |
-
|
| 45 |
-
self.features = nn.Sequential()
|
| 46 |
-
for i, channels_per_stage in enumerate(channels):
|
| 47 |
-
stage = nn.Sequential()
|
| 48 |
-
for j, out_channels in enumerate(channels_per_stage):
|
| 49 |
-
pointwise = (len(channels_per_stage) > 1) and not (((j + 1) % 2 == 1) ^ odd_pointwise)
|
| 50 |
-
stage.add_module(f"unit{j + 1}", DarkNetUnit(in_channels, out_channels, pointwise, alpha))
|
| 51 |
-
in_channels = out_channels
|
| 52 |
-
if i != len(channels) - 1:
|
| 53 |
-
stage.add_module(f"pool{i + 1}", nn.MaxPool2d(kernel_size=2, stride=2))
|
| 54 |
-
self.features.add_module(f"stage{i + 1}", stage)
|
| 55 |
-
|
| 56 |
-
final_feature_map_size = image_size // (2 ** (len(channels) - 1))
|
| 57 |
-
|
| 58 |
-
self.output = nn.Sequential(
|
| 59 |
-
nn.Conv2d(in_channels=in_channels, out_channels=num_classes, kernel_size=1),
|
| 60 |
-
nn.LeakyReLU(negative_slope=alpha, inplace=True),
|
| 61 |
-
nn.AdaptiveAvgPool2d(output_size=(1, 1))
|
| 62 |
-
)
|
| 63 |
-
|
| 64 |
-
self._initialize_weights()
|
| 65 |
-
|
| 66 |
-
def _initialize_weights(self):
|
| 67 |
-
for module in self.modules():
|
| 68 |
-
if isinstance(module, nn.Conv2d):
|
| 69 |
-
nn.init.kaiming_uniform_(module.weight, mode='fan_in', nonlinearity='leaky_relu')
|
| 70 |
-
if module.bias is not None:
|
| 71 |
-
nn.init.constant_(module.bias, 0)
|
| 72 |
-
|
| 73 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 74 |
-
x = self.features(x)
|
| 75 |
-
x = self.output(x)
|
| 76 |
-
x = x.view(x.size(0), -1)
|
| 77 |
-
return x
|
| 78 |
-
|
| 79 |
-
def train_setup(self, prm: dict):
|
| 80 |
-
self.to(self.device)
|
| 81 |
-
learning_rate = float(prm.get("lr", 0.01))
|
| 82 |
-
momentum = float(prm.get("momentum", 0.9))
|
| 83 |
-
self.criteria = nn.CrossEntropyLoss()
|
| 84 |
-
self.optimizer = optim.SGD(self.parameters(), lr=learning_rate, momentum=momentum)
|
| 85 |
-
self.to(self.device)
|
| 86 |
-
|
| 87 |
-
def learn(self, train_data: torch.utils.data.DataLoader):
|
| 88 |
-
self.train()
|
| 89 |
-
for inputs, targets in train_data:
|
| 90 |
-
inputs, targets = inputs.to(next(self.parameters()).device), targets.to(next(self.parameters()).device)
|
| 91 |
-
self.optimizer.zero_grad()
|
| 92 |
-
outputs = self(inputs)
|
| 93 |
-
loss = self.criteria(outputs, targets)
|
| 94 |
-
loss.backward()
|
| 95 |
-
self.optimizer.step()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DeepLabV3-1.py
DELETED
|
@@ -1,382 +0,0 @@
|
|
| 1 |
-
from collections import OrderedDict
|
| 2 |
-
from typing import Callable, Dict, List, Optional, Sequence, Type, Union
|
| 3 |
-
|
| 4 |
-
import torch
|
| 5 |
-
import torch.nn.functional as F
|
| 6 |
-
from torch import nn, Tensor
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
class DeepLabHead(nn.Sequential):
|
| 10 |
-
def __init__(self, in_channels: int, num_classes: int = 100, atrous_rates: Sequence[int] = (12, 24, 36)) -> None:
|
| 11 |
-
super(DeepLabHead, self).__init__(
|
| 12 |
-
ASPP(in_channels, atrous_rates),
|
| 13 |
-
nn.Conv2d(256, 256, 3, padding=1, bias=False),
|
| 14 |
-
nn.BatchNorm2d(256),
|
| 15 |
-
nn.ReLU(True),
|
| 16 |
-
nn.Conv2d(256, num_classes, 1),
|
| 17 |
-
)
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
class ASPPConv(nn.Sequential):
|
| 21 |
-
def __init__(self, in_channels: int, out_channels: int, dilation: int) -> None:
|
| 22 |
-
modules = [
|
| 23 |
-
nn.Conv2d(in_channels, out_channels, 3, padding=dilation, dilation=dilation, bias=False),
|
| 24 |
-
nn.BatchNorm2d(out_channels),
|
| 25 |
-
nn.ReLU(True),
|
| 26 |
-
]
|
| 27 |
-
super(ASPPConv, self).__init__(*modules)
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
class ASPPPooling(nn.Sequential):
|
| 31 |
-
def __init__(self, in_channels: int, out_channels: int) -> None:
|
| 32 |
-
super(ASPPPooling, self).__init__(
|
| 33 |
-
nn.AdaptiveAvgPool2d(1),
|
| 34 |
-
nn.Conv2d(in_channels, out_channels, 1, bias=False),
|
| 35 |
-
nn.BatchNorm2d(out_channels),
|
| 36 |
-
nn.ReLU(),
|
| 37 |
-
)
|
| 38 |
-
|
| 39 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 40 |
-
size = x.shape[-2:]
|
| 41 |
-
for mod in self:
|
| 42 |
-
x = mod(x)
|
| 43 |
-
return F.interpolate(x, size=size, mode="bilinear", align_corners=False)
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
class ASPP(nn.Module):
|
| 47 |
-
def __init__(self, in_channels: int, atrous_rates: Sequence[int], out_channels: int = 256) -> None:
|
| 48 |
-
super(ASPP, self).__init__()
|
| 49 |
-
modules = []
|
| 50 |
-
modules.append(
|
| 51 |
-
nn.Sequential(nn.Conv2d(in_channels, out_channels, 1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU())
|
| 52 |
-
)
|
| 53 |
-
|
| 54 |
-
rates = tuple(atrous_rates)
|
| 55 |
-
for rate in rates:
|
| 56 |
-
modules.append(ASPPConv(in_channels, out_channels, rate))
|
| 57 |
-
|
| 58 |
-
modules.append(ASPPPooling(in_channels, out_channels))
|
| 59 |
-
|
| 60 |
-
self.convs = nn.ModuleList(modules)
|
| 61 |
-
|
| 62 |
-
self.project = nn.Sequential(
|
| 63 |
-
nn.Conv2d(len(self.convs) * out_channels, out_channels, 1, bias=False),
|
| 64 |
-
nn.BatchNorm2d(out_channels),
|
| 65 |
-
nn.ReLU(),
|
| 66 |
-
nn.Dropout(0.5),
|
| 67 |
-
)
|
| 68 |
-
|
| 69 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 70 |
-
_res = []
|
| 71 |
-
for conv in self.convs:
|
| 72 |
-
_res.append(conv(x))
|
| 73 |
-
res = torch.cat(_res, dim=1)
|
| 74 |
-
return self.project(res)
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
class FCNHead(nn.Sequential):
|
| 78 |
-
def __init__(self, in_channels: int, channels: int) -> None:
|
| 79 |
-
inter_channels = in_channels // 4
|
| 80 |
-
layers = [
|
| 81 |
-
nn.Conv2d(in_channels, inter_channels, 3, padding=1, bias=False),
|
| 82 |
-
nn.BatchNorm2d(inter_channels),
|
| 83 |
-
nn.ReLU(),
|
| 84 |
-
nn.Dropout(0.1),
|
| 85 |
-
nn.Conv2d(inter_channels, channels, 1),
|
| 86 |
-
]
|
| 87 |
-
|
| 88 |
-
super(FCNHead, self).__init__(*layers)
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
def conv3x3(in_planes: int, out_planes: int, stride: int = 1, groups: int = 1, dilation: int = 1) -> nn.Conv2d:
|
| 92 |
-
return nn.Conv2d(
|
| 93 |
-
in_planes,
|
| 94 |
-
out_planes,
|
| 95 |
-
kernel_size=3,
|
| 96 |
-
stride=stride,
|
| 97 |
-
padding=dilation,
|
| 98 |
-
groups=groups,
|
| 99 |
-
bias=False,
|
| 100 |
-
dilation=dilation,
|
| 101 |
-
)
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
def conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d:
|
| 105 |
-
return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
class BasicBlock(nn.Module):
|
| 109 |
-
expansion: int = 1
|
| 110 |
-
|
| 111 |
-
def __init__(
|
| 112 |
-
self,
|
| 113 |
-
inplanes: int,
|
| 114 |
-
planes: int,
|
| 115 |
-
stride: int = 1,
|
| 116 |
-
downsample: Optional[nn.Module] = None,
|
| 117 |
-
groups: int = 1,
|
| 118 |
-
base_width: int = 64,
|
| 119 |
-
dilation: int = 1,
|
| 120 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 121 |
-
) -> None:
|
| 122 |
-
super().__init__()
|
| 123 |
-
if norm_layer is None:
|
| 124 |
-
norm_layer = nn.BatchNorm2d
|
| 125 |
-
if groups != 1 or base_width != 64:
|
| 126 |
-
raise ValueError("BasicBlock only supports groups=1 and base_width=64")
|
| 127 |
-
if dilation > 1:
|
| 128 |
-
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
|
| 129 |
-
self.conv1 = conv3x3(inplanes, planes, stride)
|
| 130 |
-
self.bn1 = norm_layer(planes)
|
| 131 |
-
self.relu = nn.ReLU(inplace=True)
|
| 132 |
-
self.conv2 = conv3x3(planes, planes)
|
| 133 |
-
self.bn2 = norm_layer(planes)
|
| 134 |
-
self.downsample = downsample
|
| 135 |
-
self.stride = stride
|
| 136 |
-
|
| 137 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 138 |
-
identity = x
|
| 139 |
-
|
| 140 |
-
out = self.conv1(x)
|
| 141 |
-
out = self.bn1(out)
|
| 142 |
-
out = self.relu(out)
|
| 143 |
-
|
| 144 |
-
out = self.conv2(out)
|
| 145 |
-
out = self.bn2(out)
|
| 146 |
-
|
| 147 |
-
if self.downsample is not None:
|
| 148 |
-
identity = self.downsample(x)
|
| 149 |
-
|
| 150 |
-
out += identity
|
| 151 |
-
out = self.relu(out)
|
| 152 |
-
|
| 153 |
-
return out
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
class Bottleneck(nn.Module):
|
| 157 |
-
expansion: int = 4
|
| 158 |
-
|
| 159 |
-
def __init__(
|
| 160 |
-
self,
|
| 161 |
-
inplanes: int,
|
| 162 |
-
planes: int,
|
| 163 |
-
stride: int = 1,
|
| 164 |
-
downsample: Optional[nn.Module] = None,
|
| 165 |
-
groups: int = 1,
|
| 166 |
-
base_width: int = 64,
|
| 167 |
-
dilation: int = 1,
|
| 168 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 169 |
-
) -> None:
|
| 170 |
-
super().__init__()
|
| 171 |
-
if norm_layer is None:
|
| 172 |
-
norm_layer = nn.BatchNorm2d
|
| 173 |
-
width = int(planes * (base_width / 64.0)) * groups
|
| 174 |
-
self.conv1 = conv1x1(inplanes, width)
|
| 175 |
-
self.bn1 = norm_layer(width)
|
| 176 |
-
self.conv2 = conv3x3(width, width, stride, groups, dilation)
|
| 177 |
-
self.bn2 = norm_layer(width)
|
| 178 |
-
self.conv3 = conv1x1(width, planes * self.expansion)
|
| 179 |
-
self.bn3 = norm_layer(planes * self.expansion)
|
| 180 |
-
self.relu = nn.ReLU(inplace=True)
|
| 181 |
-
self.downsample = downsample
|
| 182 |
-
self.stride = stride
|
| 183 |
-
|
| 184 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 185 |
-
identity = x
|
| 186 |
-
|
| 187 |
-
out = self.conv1(x)
|
| 188 |
-
out = self.bn1(out)
|
| 189 |
-
out = self.relu(out)
|
| 190 |
-
|
| 191 |
-
out = self.conv2(out)
|
| 192 |
-
out = self.bn2(out)
|
| 193 |
-
out = self.relu(out)
|
| 194 |
-
|
| 195 |
-
out = self.conv3(out)
|
| 196 |
-
out = self.bn3(out)
|
| 197 |
-
|
| 198 |
-
if self.downsample is not None:
|
| 199 |
-
identity = self.downsample(x)
|
| 200 |
-
|
| 201 |
-
out += identity
|
| 202 |
-
out = self.relu(out)
|
| 203 |
-
|
| 204 |
-
return out
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
class ResNet(nn.Module):
|
| 208 |
-
def __init__(
|
| 209 |
-
self,
|
| 210 |
-
channels: int,
|
| 211 |
-
block: Type[Union[BasicBlock, Bottleneck]],
|
| 212 |
-
layers: List[int],
|
| 213 |
-
num_classes: int = 1000,
|
| 214 |
-
zero_init_residual: bool = False,
|
| 215 |
-
groups: int = 1,
|
| 216 |
-
width_per_group: int = 64,
|
| 217 |
-
replace_stride_with_dilation: Optional[List[bool]] = None,
|
| 218 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 219 |
-
) -> None:
|
| 220 |
-
super(ResNet, self).__init__()
|
| 221 |
-
if norm_layer is None:
|
| 222 |
-
norm_layer = nn.BatchNorm2d
|
| 223 |
-
self._norm_layer = norm_layer
|
| 224 |
-
|
| 225 |
-
self.inplanes = 64
|
| 226 |
-
self.dilation = 1
|
| 227 |
-
if replace_stride_with_dilation is None:
|
| 228 |
-
replace_stride_with_dilation = [False, False, False]
|
| 229 |
-
if len(replace_stride_with_dilation) != 3:
|
| 230 |
-
raise ValueError(
|
| 231 |
-
"replace_stride_with_dilation should be None "
|
| 232 |
-
f"or a 3-element tuple, got {replace_stride_with_dilation}"
|
| 233 |
-
)
|
| 234 |
-
self.groups = groups
|
| 235 |
-
self.base_width = width_per_group
|
| 236 |
-
self.conv1 = nn.Conv2d(channels, self.inplanes, kernel_size=7, stride=2, padding=3, bias=False)
|
| 237 |
-
self.bn1 = norm_layer(self.inplanes)
|
| 238 |
-
self.relu = nn.ReLU(inplace=True)
|
| 239 |
-
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 240 |
-
self.layer1 = self._make_layer(block, 64, layers[0])
|
| 241 |
-
self.layer2 = self._make_layer(block, 128, layers[1], stride=2, dilate=replace_stride_with_dilation[0])
|
| 242 |
-
self.layer3 = self._make_layer(block, 256, layers[2], stride=2, dilate=replace_stride_with_dilation[1])
|
| 243 |
-
self.layer4 = self._make_layer(block, 512, layers[3], stride=2, dilate=replace_stride_with_dilation[2])
|
| 244 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 245 |
-
self.fc = nn.Linear(512 * block.expansion, num_classes)
|
| 246 |
-
|
| 247 |
-
for m in self.modules():
|
| 248 |
-
if isinstance(m, nn.Conv2d):
|
| 249 |
-
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
|
| 250 |
-
elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):
|
| 251 |
-
nn.init.constant_(m.weight, 1)
|
| 252 |
-
nn.init.constant_(m.bias, 0)
|
| 253 |
-
|
| 254 |
-
if zero_init_residual:
|
| 255 |
-
for m in self.modules():
|
| 256 |
-
if isinstance(m, Bottleneck) and m.bn3.weight is not None:
|
| 257 |
-
nn.init.constant_(m.bn3.weight, 0)
|
| 258 |
-
elif isinstance(m, BasicBlock) and m.bn2.weight is not None:
|
| 259 |
-
nn.init.constant_(m.bn2.weight, 0)
|
| 260 |
-
|
| 261 |
-
def _make_layer(
|
| 262 |
-
self,
|
| 263 |
-
block: Type[Union[BasicBlock, Bottleneck]],
|
| 264 |
-
planes: int,
|
| 265 |
-
blocks: int,
|
| 266 |
-
stride: int = 1,
|
| 267 |
-
dilate: bool = False,
|
| 268 |
-
) -> nn.Sequential:
|
| 269 |
-
norm_layer = self._norm_layer
|
| 270 |
-
downsample = None
|
| 271 |
-
previous_dilation = self.dilation
|
| 272 |
-
if dilate:
|
| 273 |
-
self.dilation *= stride
|
| 274 |
-
stride = 1
|
| 275 |
-
if stride != 1 or self.inplanes != planes * block.expansion:
|
| 276 |
-
downsample = nn.Sequential(
|
| 277 |
-
conv1x1(self.inplanes, planes * block.expansion, stride),
|
| 278 |
-
norm_layer(planes * block.expansion),
|
| 279 |
-
)
|
| 280 |
-
|
| 281 |
-
layers = []
|
| 282 |
-
layers.append(
|
| 283 |
-
block(
|
| 284 |
-
self.inplanes, planes, stride, downsample, self.groups, self.base_width, previous_dilation, norm_layer
|
| 285 |
-
)
|
| 286 |
-
)
|
| 287 |
-
self.inplanes = planes * block.expansion
|
| 288 |
-
for _ in range(1, blocks):
|
| 289 |
-
layers.append(
|
| 290 |
-
block(
|
| 291 |
-
self.inplanes,
|
| 292 |
-
planes,
|
| 293 |
-
groups=self.groups,
|
| 294 |
-
base_width=self.base_width,
|
| 295 |
-
dilation=self.dilation,
|
| 296 |
-
norm_layer=norm_layer,
|
| 297 |
-
)
|
| 298 |
-
)
|
| 299 |
-
|
| 300 |
-
return nn.Sequential(*layers)
|
| 301 |
-
|
| 302 |
-
def _forward_impl(self, x: Tensor) -> Tensor:
|
| 303 |
-
x = self.conv1(x)
|
| 304 |
-
x = self.bn1(x)
|
| 305 |
-
x = self.relu(x)
|
| 306 |
-
x = self.maxpool(x)
|
| 307 |
-
|
| 308 |
-
x = self.layer1(x)
|
| 309 |
-
x = self.layer2(x)
|
| 310 |
-
x = self.layer3(x)
|
| 311 |
-
x = self.layer4(x)
|
| 312 |
-
|
| 313 |
-
x = self.avgpool(x)
|
| 314 |
-
x = torch.flatten(x, 1)
|
| 315 |
-
x = self.fc(x)
|
| 316 |
-
|
| 317 |
-
return x
|
| 318 |
-
|
| 319 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 320 |
-
return self._forward_impl(x)
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
def supported_hyperparameters():
|
| 324 |
-
return {'lr', 'momentum'}
|
| 325 |
-
|
| 326 |
-
|
| 327 |
-
class Net(nn.Module):
|
| 328 |
-
|
| 329 |
-
def train_setup(self, prm):
|
| 330 |
-
self.to(self.device)
|
| 331 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=-1).to(self.device),)
|
| 332 |
-
params_list = [{'params': self.backbone.parameters(), 'lr': prm['lr']}]
|
| 333 |
-
for module in self.exclusive:
|
| 334 |
-
params_list.append({'params': getattr(self, module).parameters(), 'lr': prm['lr'] * 10})
|
| 335 |
-
self.optimizer = torch.optim.SGD(params_list, lr=prm['lr'], momentum=prm['momentum'])
|
| 336 |
-
|
| 337 |
-
def learn(self, train_data):
|
| 338 |
-
for inputs, labels in train_data:
|
| 339 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 340 |
-
self.optimizer.zero_grad()
|
| 341 |
-
outputs = self(inputs)
|
| 342 |
-
loss = self.criteria[0](outputs, labels)
|
| 343 |
-
loss.backward()
|
| 344 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 345 |
-
self.optimizer.step()
|
| 346 |
-
|
| 347 |
-
__constants__ = ["aux_classifier"]
|
| 348 |
-
|
| 349 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 350 |
-
super(Net, self).__init__()
|
| 351 |
-
self.device = device
|
| 352 |
-
num_classes = out_shape[0]
|
| 353 |
-
self.backbone: nn.Module = ResNet(in_shape[1], Bottleneck, [3, 4, 6, 3], num_classes=100, replace_stride_with_dilation=[False, True, True])
|
| 354 |
-
self.classifier: nn.Module = DeepLabHead(2048, num_classes)
|
| 355 |
-
self.aux_classifier: Optional[nn.Module] = None
|
| 356 |
-
self.__setattr__('exclusive', ['classifier'] if self.aux_classifier == None else ['classifier', 'aux_classifier'])
|
| 357 |
-
|
| 358 |
-
def forward(self, x: Tensor) -> Union[Dict[str, Tensor], Tensor]:
|
| 359 |
-
input_shape = x.shape[-2:]
|
| 360 |
-
c3, c4 = self.backbone_fw(x)
|
| 361 |
-
x = self.classifier(c4)
|
| 362 |
-
x = F.interpolate(x, size=input_shape, mode="bilinear", align_corners=False)
|
| 363 |
-
|
| 364 |
-
if self.aux_classifier is not None:
|
| 365 |
-
result = OrderedDict()
|
| 366 |
-
result["out"] = x
|
| 367 |
-
x = self.aux_classifier(c3)
|
| 368 |
-
x = F.interpolate(x, size=input_shape, mode="bilinear", align_corners=False)
|
| 369 |
-
result["aux"] = x
|
| 370 |
-
return result
|
| 371 |
-
return x
|
| 372 |
-
|
| 373 |
-
def backbone_fw(self, x):
|
| 374 |
-
x = self.backbone.conv1(x)
|
| 375 |
-
x = self.backbone.bn1(x)
|
| 376 |
-
x = self.backbone.relu(x)
|
| 377 |
-
x = self.backbone.maxpool(x)
|
| 378 |
-
c1 = self.backbone.layer1(x)
|
| 379 |
-
c2 = self.backbone.layer2(c1)
|
| 380 |
-
c3 = self.backbone.layer3(c2)
|
| 381 |
-
c4 = self.backbone.layer4(c3)
|
| 382 |
-
return c3, c4
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DeepLabV3-2.py
DELETED
|
@@ -1,382 +0,0 @@
|
|
| 1 |
-
from collections import OrderedDict
|
| 2 |
-
from typing import Callable, Dict, List, Optional, Sequence, Type, Union
|
| 3 |
-
|
| 4 |
-
import torch
|
| 5 |
-
import torch.nn.functional as F
|
| 6 |
-
from torch import nn, Tensor
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
class DeepLabHead(nn.Sequential):
|
| 10 |
-
def __init__(self, in_channels: int, num_classes: int = 100, atrous_rates: Sequence[int] = (12, 24, 36)) -> None:
|
| 11 |
-
super(DeepLabHead, self).__init__(
|
| 12 |
-
ASPP(in_channels, atrous_rates),
|
| 13 |
-
nn.Conv2d(256, 256, 3, padding=1, bias=False),
|
| 14 |
-
nn.BatchNorm2d(256),
|
| 15 |
-
nn.ReLU(True),
|
| 16 |
-
nn.Conv2d(256, num_classes, 1),
|
| 17 |
-
)
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
class ASPPConv(nn.Sequential):
|
| 21 |
-
def __init__(self, in_channels: int, out_channels: int, dilation: int) -> None:
|
| 22 |
-
modules = [
|
| 23 |
-
nn.Conv2d(in_channels, out_channels, 3, padding=dilation, dilation=dilation, bias=False),
|
| 24 |
-
nn.BatchNorm2d(out_channels),
|
| 25 |
-
nn.ReLU(True),
|
| 26 |
-
]
|
| 27 |
-
super(ASPPConv, self).__init__(*modules)
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
class ASPPPooling(nn.Sequential):
|
| 31 |
-
def __init__(self, in_channels: int, out_channels: int) -> None:
|
| 32 |
-
super(ASPPPooling, self).__init__(
|
| 33 |
-
nn.AdaptiveAvgPool2d(1),
|
| 34 |
-
nn.Conv2d(in_channels, out_channels, 1, bias=False),
|
| 35 |
-
nn.BatchNorm2d(out_channels),
|
| 36 |
-
nn.ReLU(),
|
| 37 |
-
)
|
| 38 |
-
|
| 39 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 40 |
-
size = x.shape[-2:]
|
| 41 |
-
for mod in self:
|
| 42 |
-
x = mod(x)
|
| 43 |
-
return F.interpolate(x, size=size, mode="bilinear", align_corners=False)
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
class ASPP(nn.Module):
|
| 47 |
-
def __init__(self, in_channels: int, atrous_rates: Sequence[int], out_channels: int = 256) -> None:
|
| 48 |
-
super(ASPP, self).__init__()
|
| 49 |
-
modules = []
|
| 50 |
-
modules.append(
|
| 51 |
-
nn.Sequential(nn.Conv2d(in_channels, out_channels, 1, bias=False), nn.BatchNorm2d(out_channels), nn.ReLU())
|
| 52 |
-
)
|
| 53 |
-
|
| 54 |
-
rates = tuple(atrous_rates)
|
| 55 |
-
for rate in rates:
|
| 56 |
-
modules.append(ASPPConv(in_channels, out_channels, rate))
|
| 57 |
-
|
| 58 |
-
modules.append(ASPPPooling(in_channels, out_channels))
|
| 59 |
-
|
| 60 |
-
self.convs = nn.ModuleList(modules)
|
| 61 |
-
|
| 62 |
-
self.project = nn.Sequential(
|
| 63 |
-
nn.Conv2d(len(self.convs) * out_channels, out_channels, 1, bias=False),
|
| 64 |
-
nn.BatchNorm2d(out_channels),
|
| 65 |
-
nn.ReLU(),
|
| 66 |
-
nn.Dropout(0.5),
|
| 67 |
-
)
|
| 68 |
-
|
| 69 |
-
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 70 |
-
_res = []
|
| 71 |
-
for conv in self.convs:
|
| 72 |
-
_res.append(conv(x))
|
| 73 |
-
res = torch.cat(_res, dim=1)
|
| 74 |
-
return self.project(res)
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
class FCNHead(nn.Sequential):
|
| 78 |
-
def __init__(self, in_channels: int, channels: int) -> None:
|
| 79 |
-
inter_channels = in_channels // 4
|
| 80 |
-
layers = [
|
| 81 |
-
nn.Conv2d(in_channels, inter_channels, 3, padding=1, bias=False),
|
| 82 |
-
nn.BatchNorm2d(inter_channels),
|
| 83 |
-
nn.ReLU(),
|
| 84 |
-
nn.Dropout(0.1),
|
| 85 |
-
nn.Conv2d(inter_channels, channels, 1),
|
| 86 |
-
]
|
| 87 |
-
|
| 88 |
-
super(FCNHead, self).__init__(*layers)
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
def conv3x3(in_planes: int, out_planes: int, stride: int = 1, groups: int = 1, dilation: int = 1) -> nn.Conv2d:
|
| 92 |
-
return nn.Conv2d(
|
| 93 |
-
in_planes,
|
| 94 |
-
out_planes,
|
| 95 |
-
kernel_size=3,
|
| 96 |
-
stride=stride,
|
| 97 |
-
padding=dilation,
|
| 98 |
-
groups=groups,
|
| 99 |
-
bias=False,
|
| 100 |
-
dilation=dilation,
|
| 101 |
-
)
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
def conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d:
|
| 105 |
-
return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
class BasicBlock(nn.Module):
|
| 109 |
-
expansion: int = 1
|
| 110 |
-
|
| 111 |
-
def __init__(
|
| 112 |
-
self,
|
| 113 |
-
inplanes: int,
|
| 114 |
-
planes: int,
|
| 115 |
-
stride: int = 1,
|
| 116 |
-
downsample: Optional[nn.Module] = None,
|
| 117 |
-
groups: int = 1,
|
| 118 |
-
base_width: int = 64,
|
| 119 |
-
dilation: int = 1,
|
| 120 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 121 |
-
) -> None:
|
| 122 |
-
super().__init__()
|
| 123 |
-
if norm_layer is None:
|
| 124 |
-
norm_layer = nn.BatchNorm2d
|
| 125 |
-
if groups != 1 or base_width != 64:
|
| 126 |
-
raise ValueError("BasicBlock only supports groups=1 and base_width=64")
|
| 127 |
-
if dilation > 1:
|
| 128 |
-
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
|
| 129 |
-
self.conv1 = conv3x3(inplanes, planes, stride)
|
| 130 |
-
self.bn1 = norm_layer(planes)
|
| 131 |
-
self.relu = nn.ReLU(inplace=True)
|
| 132 |
-
self.conv2 = conv3x3(planes, planes)
|
| 133 |
-
self.bn2 = norm_layer(planes)
|
| 134 |
-
self.downsample = downsample
|
| 135 |
-
self.stride = stride
|
| 136 |
-
|
| 137 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 138 |
-
identity = x
|
| 139 |
-
|
| 140 |
-
out = self.conv1(x)
|
| 141 |
-
out = self.bn1(out)
|
| 142 |
-
out = self.relu(out)
|
| 143 |
-
|
| 144 |
-
out = self.conv2(out)
|
| 145 |
-
out = self.bn2(out)
|
| 146 |
-
|
| 147 |
-
if self.downsample is not None:
|
| 148 |
-
identity = self.downsample(x)
|
| 149 |
-
|
| 150 |
-
out += identity
|
| 151 |
-
out = self.relu(out)
|
| 152 |
-
|
| 153 |
-
return out
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
class Bottleneck(nn.Module):
|
| 157 |
-
expansion: int = 4
|
| 158 |
-
|
| 159 |
-
def __init__(
|
| 160 |
-
self,
|
| 161 |
-
inplanes: int,
|
| 162 |
-
planes: int,
|
| 163 |
-
stride: int = 1,
|
| 164 |
-
downsample: Optional[nn.Module] = None,
|
| 165 |
-
groups: int = 1,
|
| 166 |
-
base_width: int = 64,
|
| 167 |
-
dilation: int = 1,
|
| 168 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 169 |
-
) -> None:
|
| 170 |
-
super().__init__()
|
| 171 |
-
if norm_layer is None:
|
| 172 |
-
norm_layer = nn.BatchNorm2d
|
| 173 |
-
width = int(planes * (base_width / 64.0)) * groups
|
| 174 |
-
self.conv1 = conv1x1(inplanes, width)
|
| 175 |
-
self.bn1 = norm_layer(width)
|
| 176 |
-
self.conv2 = conv3x3(width, width, stride, groups, dilation)
|
| 177 |
-
self.bn2 = norm_layer(width)
|
| 178 |
-
self.conv3 = conv1x1(width, planes * self.expansion)
|
| 179 |
-
self.bn3 = norm_layer(planes * self.expansion)
|
| 180 |
-
self.relu = nn.ReLU(inplace=True)
|
| 181 |
-
self.downsample = downsample
|
| 182 |
-
self.stride = stride
|
| 183 |
-
|
| 184 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 185 |
-
identity = x
|
| 186 |
-
|
| 187 |
-
out = self.conv1(x)
|
| 188 |
-
out = self.bn1(out)
|
| 189 |
-
out = self.relu(out)
|
| 190 |
-
|
| 191 |
-
out = self.conv2(out)
|
| 192 |
-
out = self.bn2(out)
|
| 193 |
-
out = self.relu(out)
|
| 194 |
-
|
| 195 |
-
out = self.conv3(out)
|
| 196 |
-
out = self.bn3(out)
|
| 197 |
-
|
| 198 |
-
if self.downsample is not None:
|
| 199 |
-
identity = self.downsample(x)
|
| 200 |
-
|
| 201 |
-
out += identity
|
| 202 |
-
out = self.relu(out)
|
| 203 |
-
|
| 204 |
-
return out
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
class ResNet(nn.Module):
|
| 208 |
-
def __init__(
|
| 209 |
-
self,
|
| 210 |
-
channels: int,
|
| 211 |
-
block: Type[Union[BasicBlock, Bottleneck]],
|
| 212 |
-
layers: List[int],
|
| 213 |
-
num_classes: int = 1000,
|
| 214 |
-
zero_init_residual: bool = False,
|
| 215 |
-
groups: int = 1,
|
| 216 |
-
width_per_group: int = 64,
|
| 217 |
-
replace_stride_with_dilation: Optional[List[bool]] = None,
|
| 218 |
-
norm_layer: Optional[Callable[..., nn.Module]] = None,
|
| 219 |
-
) -> None:
|
| 220 |
-
super(ResNet, self).__init__()
|
| 221 |
-
if norm_layer is None:
|
| 222 |
-
norm_layer = nn.BatchNorm2d
|
| 223 |
-
self._norm_layer = norm_layer
|
| 224 |
-
|
| 225 |
-
self.inplanes = 64
|
| 226 |
-
self.dilation = 1
|
| 227 |
-
if replace_stride_with_dilation is None:
|
| 228 |
-
replace_stride_with_dilation = [False, False, False]
|
| 229 |
-
if len(replace_stride_with_dilation) != 3:
|
| 230 |
-
raise ValueError(
|
| 231 |
-
"replace_stride_with_dilation should be None "
|
| 232 |
-
f"or a 3-element tuple, got {replace_stride_with_dilation}"
|
| 233 |
-
)
|
| 234 |
-
self.groups = groups
|
| 235 |
-
self.base_width = width_per_group
|
| 236 |
-
self.conv1 = nn.Conv2d(channels, self.inplanes, kernel_size=7, stride=2, padding=3, bias=False)
|
| 237 |
-
self.bn1 = norm_layer(self.inplanes)
|
| 238 |
-
self.relu = nn.ReLU(inplace=True)
|
| 239 |
-
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
| 240 |
-
self.layer1 = self._make_layer(block, 64, layers[0])
|
| 241 |
-
self.layer2 = self._make_layer(block, 128, layers[1], stride=2, dilate=replace_stride_with_dilation[0])
|
| 242 |
-
self.layer3 = self._make_layer(block, 256, layers[2], stride=2, dilate=replace_stride_with_dilation[1])
|
| 243 |
-
self.layer4 = self._make_layer(block, 512, layers[3], stride=2, dilate=replace_stride_with_dilation[2])
|
| 244 |
-
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
| 245 |
-
self.fc = nn.Linear(512 * block.expansion, num_classes)
|
| 246 |
-
|
| 247 |
-
for m in self.modules():
|
| 248 |
-
if isinstance(m, nn.Conv2d):
|
| 249 |
-
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
|
| 250 |
-
elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):
|
| 251 |
-
nn.init.constant_(m.weight, 1)
|
| 252 |
-
nn.init.constant_(m.bias, 0)
|
| 253 |
-
|
| 254 |
-
if zero_init_residual:
|
| 255 |
-
for m in self.modules():
|
| 256 |
-
if isinstance(m, Bottleneck) and m.bn3.weight is not None:
|
| 257 |
-
nn.init.constant_(m.bn3.weight, 0)
|
| 258 |
-
elif isinstance(m, BasicBlock) and m.bn2.weight is not None:
|
| 259 |
-
nn.init.constant_(m.bn2.weight, 0)
|
| 260 |
-
|
| 261 |
-
def _make_layer(
|
| 262 |
-
self,
|
| 263 |
-
block: Type[Union[BasicBlock, Bottleneck]],
|
| 264 |
-
planes: int,
|
| 265 |
-
blocks: int,
|
| 266 |
-
stride: int = 1,
|
| 267 |
-
dilate: bool = False,
|
| 268 |
-
) -> nn.Sequential:
|
| 269 |
-
norm_layer = self._norm_layer
|
| 270 |
-
downsample = None
|
| 271 |
-
previous_dilation = self.dilation
|
| 272 |
-
if dilate:
|
| 273 |
-
self.dilation *= stride
|
| 274 |
-
stride = 1
|
| 275 |
-
if stride != 1 or self.inplanes != planes * block.expansion:
|
| 276 |
-
downsample = nn.Sequential(
|
| 277 |
-
conv1x1(self.inplanes, planes * block.expansion, stride),
|
| 278 |
-
norm_layer(planes * block.expansion),
|
| 279 |
-
)
|
| 280 |
-
|
| 281 |
-
layers = []
|
| 282 |
-
layers.append(
|
| 283 |
-
block(
|
| 284 |
-
self.inplanes, planes, stride, downsample, self.groups, self.base_width, previous_dilation, norm_layer
|
| 285 |
-
)
|
| 286 |
-
)
|
| 287 |
-
self.inplanes = planes * block.expansion
|
| 288 |
-
for _ in range(1, blocks):
|
| 289 |
-
layers.append(
|
| 290 |
-
block(
|
| 291 |
-
self.inplanes,
|
| 292 |
-
planes,
|
| 293 |
-
groups=self.groups,
|
| 294 |
-
base_width=self.base_width,
|
| 295 |
-
dilation=self.dilation,
|
| 296 |
-
norm_layer=norm_layer,
|
| 297 |
-
)
|
| 298 |
-
)
|
| 299 |
-
|
| 300 |
-
return nn.Sequential(*layers)
|
| 301 |
-
|
| 302 |
-
def _forward_impl(self, x: Tensor) -> Tensor:
|
| 303 |
-
x = self.conv1(x)
|
| 304 |
-
x = self.bn1(x)
|
| 305 |
-
x = self.relu(x)
|
| 306 |
-
x = self.maxpool(x)
|
| 307 |
-
|
| 308 |
-
x = self.layer1(x)
|
| 309 |
-
x = self.layer2(x)
|
| 310 |
-
x = self.layer3(x)
|
| 311 |
-
x = self.layer4(x)
|
| 312 |
-
|
| 313 |
-
x = self.avgpool(x)
|
| 314 |
-
x = torch.flatten(x, 1)
|
| 315 |
-
x = self.fc(x)
|
| 316 |
-
|
| 317 |
-
return x
|
| 318 |
-
|
| 319 |
-
def forward(self, x: Tensor) -> Tensor:
|
| 320 |
-
return self._forward_impl(x)
|
| 321 |
-
|
| 322 |
-
|
| 323 |
-
def supported_hyperparameters():
|
| 324 |
-
return {'lr', 'momentum'}
|
| 325 |
-
|
| 326 |
-
|
| 327 |
-
class Net(nn.Module):
|
| 328 |
-
|
| 329 |
-
def train_setup(self, prm):
|
| 330 |
-
self.to(self.device)
|
| 331 |
-
self.criteria = (nn.CrossEntropyLoss(ignore_index=-1).to(self.device),)
|
| 332 |
-
params_list = [{'params': self.backbone.parameters(), 'lr': prm['lr']}]
|
| 333 |
-
for module in self.exclusive:
|
| 334 |
-
params_list.append({'params': getattr(self, module).parameters(), 'lr': prm['lr'] * 10})
|
| 335 |
-
self.optimizer = torch.optim.SGD(params_list, lr=prm['lr'], momentum=prm['momentum'])
|
| 336 |
-
|
| 337 |
-
def learn(self, train_data):
|
| 338 |
-
for inputs, labels in train_data:
|
| 339 |
-
inputs, labels = inputs.to(self.device), labels.to(self.device)
|
| 340 |
-
self.optimizer.zero_grad()
|
| 341 |
-
outputs = self(inputs)
|
| 342 |
-
loss = self.criteria[0](outputs, labels)
|
| 343 |
-
loss.backward()
|
| 344 |
-
nn.utils.clip_grad_norm_(self.parameters(), 3)
|
| 345 |
-
self.optimizer.step()
|
| 346 |
-
|
| 347 |
-
__constants__ = ["aux_classifier"]
|
| 348 |
-
|
| 349 |
-
def __init__(self, in_shape: tuple, out_shape: tuple, prm: dict, device: torch.device) -> None:
|
| 350 |
-
super(Net, self).__init__()
|
| 351 |
-
self.device = device
|
| 352 |
-
num_classes = out_shape[0]
|
| 353 |
-
self.backbone: nn.Module = ResNet(in_shape[1], Bottleneck, [3, 4, 23, 3], num_classes=100, replace_stride_with_dilation=[False, True, True])
|
| 354 |
-
self.classifier: nn.Module = DeepLabHead(2048, num_classes)
|
| 355 |
-
self.aux_classifier: Optional[nn.Module] = None
|
| 356 |
-
self.__setattr__('exclusive', ['classifier'] if self.aux_classifier == None else ['classifier', 'aux_classifier'])
|
| 357 |
-
|
| 358 |
-
def forward(self, x: Tensor) -> Union[Dict[str, Tensor], Tensor]:
|
| 359 |
-
input_shape = x.shape[-2:]
|
| 360 |
-
c3, c4 = self.backbone_fw(x)
|
| 361 |
-
x = self.classifier(c4)
|
| 362 |
-
x = F.interpolate(x, size=input_shape, mode="bilinear", align_corners=False)
|
| 363 |
-
|
| 364 |
-
if self.aux_classifier is not None:
|
| 365 |
-
result = OrderedDict()
|
| 366 |
-
result["out"] = x
|
| 367 |
-
x = self.aux_classifier(c3)
|
| 368 |
-
x = F.interpolate(x, size=input_shape, mode="bilinear", align_corners=False)
|
| 369 |
-
result["aux"] = x
|
| 370 |
-
return result
|
| 371 |
-
return x
|
| 372 |
-
|
| 373 |
-
def backbone_fw(self, x):
|
| 374 |
-
x = self.backbone.conv1(x)
|
| 375 |
-
x = self.backbone.bn1(x)
|
| 376 |
-
x = self.backbone.relu(x)
|
| 377 |
-
x = self.backbone.maxpool(x)
|
| 378 |
-
c1 = self.backbone.layer1(x)
|
| 379 |
-
c2 = self.backbone.layer2(c1)
|
| 380 |
-
c3 = self.backbone.layer3(c2)
|
| 381 |
-
c4 = self.backbone.layer4(c3)
|
| 382 |
-
return c3, c4
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
test/nn/DenoiseUNet.py
DELETED
|
@@ -1,135 +0,0 @@
|
|
| 1 |
-
import torch
|
| 2 |
-
import torch.nn as nn
|
| 3 |
-
import torch.nn.functional as F
|
| 4 |
-
import torch.optim as optim
|
| 5 |
-
|
| 6 |
-
def supported_hyperparameters():
|
| 7 |
-
return {'lr'}
|
| 8 |
-
|
| 9 |
-
class DoubleConv(nn.Module):
|
| 10 |
-
"""(convolution => [BN] => ReLU) * 2"""
|
| 11 |
-
def __init__(self, in_channels, out_channels):
|
| 12 |
-
super().__init__()
|
| 13 |
-
self.double_conv = nn.Sequential(
|
| 14 |
-
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, bias=False),
|
| 15 |
-
nn.BatchNorm2d(out_channels),
|
| 16 |
-
nn.ReLU(inplace=True),
|
| 17 |
-
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False),
|
| 18 |
-
nn.BatchNorm2d(out_channels),
|
| 19 |
-
nn.ReLU(inplace=True)
|
| 20 |
-
)
|
| 21 |
-
|
| 22 |
-
def forward(self, x):
|
| 23 |
-
return self.double_conv(x)
|
| 24 |
-
|
| 25 |
-
class Down(nn.Module):
|
| 26 |
-
"""Downscaling with maxpool then double conv"""
|
| 27 |
-
def __init__(self, in_channels, out_channels):
|
| 28 |
-
super().__init__()
|
| 29 |
-
self.maxpool_conv = nn.Sequential(
|
| 30 |
-
nn.MaxPool2d(2),
|
| 31 |
-
DoubleConv(in_channels, out_channels)
|
| 32 |
-
)
|
| 33 |
-
|
| 34 |
-
def forward(self, x):
|
| 35 |
-
return self.maxpool_conv(x)
|
| 36 |
-
|
| 37 |
-
class Up(nn.Module):
|
| 38 |
-
"""Upscaling then double conv"""
|
| 39 |
-
def __init__(self, in_channels, out_channels, bilinear=True):
|
| 40 |
-
super().__init__()
|
| 41 |
-
|
| 42 |
-
if bilinear:
|
| 43 |
-
self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
|
| 44 |
-
self.conv = DoubleConv(in_channels + (in_channels // 2), out_channels)
|
| 45 |
-
else:
|
| 46 |
-
self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)
|
| 47 |
-
self.conv = DoubleConv(in_channels, out_channels)
|
| 48 |
-
|
| 49 |
-
def forward(self, x1, x2):
|
| 50 |
-
x1 = self.up(x1)
|
| 51 |
-
diffY = x2.size()[2] - x1.size()[2]
|
| 52 |
-
diffX = x2.size()[3] - x1.size()[3]
|
| 53 |
-
x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,
|
| 54 |
-
diffY // 2, diffY - diffY // 2])
|
| 55 |
-
x = torch.cat([x2, x1], dim=1)
|
| 56 |
-
return self.conv(x)
|
| 57 |
-
|
| 58 |
-
class OutConv(nn.Module):
|
| 59 |
-
def __init__(self, in_channels, out_channels):
|
| 60 |
-
super(OutConv, self).__init__()
|
| 61 |
-
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)
|
| 62 |
-
|
| 63 |
-
def forward(self, x):
|
| 64 |
-
return self.conv(x)
|
| 65 |
-
|
| 66 |
-
class Net(nn.Module):
|
| 67 |
-
"""
|
| 68 |
-
Residual U-Net with Safety Clamping
|
| 69 |
-
"""
|
| 70 |
-
def __init__(self, in_shape, out_shape, prm, device):
|
| 71 |
-
super(Net, self).__init__()
|
| 72 |
-
self.device = device
|
| 73 |
-
n_channels = in_shape[1]
|
| 74 |
-
|
| 75 |
-
self.inc = DoubleConv(n_channels, 64)
|
| 76 |
-
self.down1 = Down(64, 128)
|
| 77 |
-
self.down2 = Down(128, 256)
|
| 78 |
-
self.down3 = Down(256, 512)
|
| 79 |
-
self.down4 = Down(512, 1024)
|
| 80 |
-
|
| 81 |
-
self.up1 = Up(1024, 512)
|
| 82 |
-
self.up2 = Up(512, 256)
|
| 83 |
-
self.up3 = Up(256, 128)
|
| 84 |
-
self.up4 = Up(128, 64)
|
| 85 |
-
self.outc = OutConv(64, n_channels)
|
| 86 |
-
|
| 87 |
-
self.to(self.device)
|
| 88 |
-
self._initialize_optimizer(prm)
|
| 89 |
-
self.criterion = nn.MSELoss()
|
| 90 |
-
|
| 91 |
-
def _initialize_optimizer(self, prm):
|
| 92 |
-
raw_lr = prm.get('lr', 0.001)
|
| 93 |
-
self.optimizer = optim.Adam(self.parameters(), lr=raw_lr)
|
| 94 |
-
|
| 95 |
-
def forward(self, x):
|
| 96 |
-
input_img = x
|
| 97 |
-
|
| 98 |
-
x1 = self.inc(x)
|
| 99 |
-
x2 = self.down1(x1)
|
| 100 |
-
x3 = self.down2(x2)
|
| 101 |
-
x4 = self.down3(x3)
|
| 102 |
-
x5 = self.down4(x4)
|
| 103 |
-
|
| 104 |
-
dec = self.up1(x5, x4)
|
| 105 |
-
dec = self.up2(dec, x3)
|
| 106 |
-
dec = self.up3(dec, x2)
|
| 107 |
-
dec = self.up4(dec, x1)
|
| 108 |
-
|
| 109 |
-
logits = self.outc(dec)
|
| 110 |
-
|
| 111 |
-
return torch.clamp(input_img + logits, 0.0, 1.0)
|
| 112 |
-
|
| 113 |
-
def train_setup(self, prm):
|
| 114 |
-
self._initialize_optimizer(prm)
|
| 115 |
-
|
| 116 |
-
def learn(self, train_data):
|
| 117 |
-
self.train()
|
| 118 |
-
total_loss = 0.0
|
| 119 |
-
count = 0
|
| 120 |
-
|
| 121 |
-
for inputs, targets in train_data:
|
| 122 |
-
inputs = inputs.to(self.device)
|
| 123 |
-
targets = targets.to(self.device)
|
| 124 |
-
|
| 125 |
-
self.optimizer.zero_grad()
|
| 126 |
-
outputs = self(inputs)
|
| 127 |
-
|
| 128 |
-
loss = self.criterion(outputs, targets)
|
| 129 |
-
loss.backward()
|
| 130 |
-
self.optimizer.step()
|
| 131 |
-
|
| 132 |
-
total_loss += loss.item()
|
| 133 |
-
count += 1
|
| 134 |
-
|
| 135 |
-
return total_loss / count if count > 0 else 0.0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|