ABrain-One commited on
Commit
e4d07b8
·
verified ·
1 Parent(s): 94b0a10

chore: Clean test directory (remove old folders/zips, rename db)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. test/{ab.nn.db.zst → ab.nn.zst-2.2.9} +0 -0
  2. test/lemur_data_sync.zip +0 -3
  3. test/nn/AirNet-626c3eb9-5c0a-43be-bcba-745592729769.py +0 -97
  4. test/nn/AirNet-777bc9dc-e6c7-4bff-8be1-a1b02adea68f.py +0 -97
  5. test/nn/AirNext-1c889567-e226-44b8-9ced-2cbb8ad0a561.py +0 -126
  6. test/nn/AirNext-31e78268-e095-43be-bc9b-ad8b34e76201.py +0 -128
  7. test/nn/AirNext-8c916d56-4362-4ab3-8f8a-94b73fe876fb.py +0 -126
  8. test/nn/AirNext.py +0 -125
  9. test/nn/AlexNet-69c52339-4eac-45f1-bbfe-2c51949701f1.py +0 -62
  10. test/nn/AlexNet-ad69700d-0e12-458f-afad-93f03988a4e7.py +0 -64
  11. test/nn/AlexNet-bb84fa5d-5dd8-4cc0-8a30-5c41958bcd94.py +0 -63
  12. test/nn/AlexNet.py +0 -61
  13. test/nn/BagNet-001bf3e2-17c2-4fdf-948e-493677e58a3b.py +0 -134
  14. test/nn/BagNet-560341b5-15a8-4829-a8ac-ea4c7391a950.py +0 -129
  15. test/nn/BagNet-6ecd3fc7-5ce2-4876-86a6-5e8250c43d78.py +0 -129
  16. test/nn/BagNet-7e541be1-6b60-445d-bbbf-3b655eeefc9a.py +0 -129
  17. test/nn/BagNet-7ebe6562-46c6-4406-96a2-bf3914ac8516.py +0 -139
  18. test/nn/BagNet-7f792262-31cf-477e-a78a-3494c122332d.py +0 -129
  19. test/nn/BayesianNet-024b0436-9ad5-4a1f-86d0-946e577ffc2d.py +0 -244
  20. test/nn/BayesianNet-0901ac22-d7f5-4deb-94c9-970e7955bd68.py +0 -242
  21. test/nn/BayesianNet-1.py +0 -241
  22. test/nn/BayesianNet-4f11c8da-cfe1-46ba-b5d0-b5d899929a2e.py +0 -238
  23. test/nn/C10C-RESNETLSTM-6a517327bf0ef897a22186a2061e85b3.py +0 -180
  24. test/nn/C10C-RESNETLSTM-8f7ac9c241d5b9546f8cd3484e0e100b.py +0 -245
  25. test/nn/C10C-RESNETLSTM-IMG-CAP-IMPROVED.py +0 -230
  26. test/nn/C10C-ResNetTransformer-187ccbee8050ac295637ecedecb4da1e.py +0 -193
  27. test/nn/C5C-RESNETLSTM-4.py +0 -222
  28. test/nn/C5C-RESNETLSTM-c42512d71480c8ef10f31e3e6c33bbdf.py +0 -150
  29. test/nn/C5C-ResNetTransformer-83fb6b6bb7c76b742ad0713d29463514.py +0 -181
  30. test/nn/C8C-ResNetTransformer-7730b6eb6979d27e2e1bbc7d05255dff.py +0 -239
  31. test/nn/ComplexNet.py +0 -295
  32. test/nn/ConditionalDiffusion.py +0 -230
  33. test/nn/ConditionalGAN.py +0 -278
  34. test/nn/ConditionalVAE3.py +0 -213
  35. test/nn/ConditionalVAE4.py +0 -268
  36. test/nn/ConvNeXt-dda5bf19-9ac1-460b-9bfd-735eec2f4904.py +0 -172
  37. test/nn/DPN107.py +0 -92
  38. test/nn/DPN131-8e6e495b-85cb-4a71-8b91-6d89372e0a0c.py +0 -86
  39. test/nn/DPN131-c53a40b8-b874-4c8b-999b-0944a1173a46.py +0 -86
  40. test/nn/DPN131-e8980802-6b89-4170-8608-327297706df0.py +0 -86
  41. test/nn/DPN131.py +0 -85
  42. test/nn/DPN68-9693aa0b-80bf-4393-a9e0-dd985a5ab128.py +0 -83
  43. test/nn/DPN68-c9cdb196-7596-4368-974a-56edf8b10381.py +0 -83
  44. test/nn/DarkNet-11e8caec-5e73-461e-a101-3aa39dfec644.py +0 -96
  45. test/nn/DarkNet-51277b91-c3a1-4669-9fb5-849ea97bd1b4.py +0 -95
  46. test/nn/DarkNet-d434ba1c-25ea-4160-a41d-4c477dba7bc0.py +0 -96
  47. test/nn/DarkNet.py +0 -95
  48. test/nn/DeepLabV3-1.py +0 -382
  49. test/nn/DeepLabV3-2.py +0 -382
  50. 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