CoolFace
Apppublic

fish0610/TW_invoice_system

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
unet_model.py86 linesDownload Raw Back to src
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)