fish0610/TW_invoice_system
0
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4
5
6class DoubleConv(nn.Module):
7 def __init__(self, in_ch, out_ch):
8 super().__init__()
9 self.net = nn.Sequential(
10 nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
11 nn.BatchNorm2d(out_ch),
12 nn.ReLU(inplace=True),
13
14 nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),
15 nn.BatchNorm2d(out_ch),
16 nn.ReLU(inplace=True),
17 )
18
19 def forward(self, x):
20 return self.net(x)
21
22
23class UNet(nn.Module):
24 def __init__(self, n_channels=3, n_classes=3):
25 super().__init__()
26 self.n_channels = n_channels
27 self.n_classes = n_classes
28
29 self.down1 = DoubleConv(n_channels, 64)
30 self.down2 = DoubleConv(64, 128)
31 self.down3 = DoubleConv(128, 256)
32 self.down4 = DoubleConv(256, 512)
33
34 self.pool = nn.MaxPool2d(2)
35
36 self.bottleneck = DoubleConv(512, 1024)
37
38 self.up4 = nn.ConvTranspose2d(1024, 512, 2, stride=2)
39 self.conv4 = DoubleConv(1024, 512)
40
41 self.up3 = nn.ConvTranspose2d(512, 256, 2, stride=2)
42 self.conv3 = DoubleConv(512, 256)
43
44 self.up2 = nn.ConvTranspose2d(256, 128, 2, stride=2)
45 self.conv2 = DoubleConv(256, 128)
46
47 self.up1 = nn.ConvTranspose2d(128, 64, 2, stride=2)
48 self.conv1 = DoubleConv(128, 64)
49
50 self.out_conv = nn.Conv2d(64, n_classes, kernel_size=1)
51
52 # 初始化 bias(小目標專案必備)
53 nn.init.constant_(self.out_conv.bias, -4)
54
55 def forward(self, x):
56 c1 = self.down1(x)
57 p1 = self.pool(c1)
58
59 c2 = self.down2(p1)
60 p2 = self.pool(c2)
61
62 c3 = self.down3(p2)
63 p3 = self.pool(c3)
64
65 c4 = self.down4(p3)
66 p4 = self.pool(c4)
67
68 bn = self.bottleneck(p4)
69
70 u4 = self.up4(bn)
71 u4 = torch.cat([u4, c4], dim=1)
72 c5 = self.conv4(u4)
73
74 u3 = self.up3(c5)
75 u3 = torch.cat([u3, c3], dim=1)
76 c6 = self.conv3(u3)
77
78 u2 = self.up2(c6)
79 u2 = torch.cat([u2, c2], dim=1)
80 c7 = self.conv2(u2)
81
82 u1 = self.up1(c7)
83 u1 = torch.cat([u1, c1], dim=1)
84 c8 = self.conv1(u1)
85
86 return self.out_conv(c8)