import torch import torch.nn as nn class CookieNet(nn.Module): def __init__(self): super(CookieNet, self).__init__() # 特徵提取層 (CNN 200x112 -> 25x14) self.features = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, stride=1, padding=1), nn.ReLU(), nn.MaxPool2d(2, 2), nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1), nn.ReLU(), nn.MaxPool2d(2, 2), nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1), nn.ReLU(), nn.MaxPool2d(2, 2), ) # 分類層 (4 類: 0:None, 1:Jump, 2:Slide, 3:Enter) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 14 * 25, 512), nn.ReLU(), nn.Dropout(0.5), nn.Linear(512, 4) ) def forward(self, x): x = self.features(x) x = self.classifier(x) return x