mirror of
https://github.com/wassname/Deep-reinforcement-learning-with-pytorch.git
synced 2026-09-11 11:52:53 +08:00
Create resnet.py
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
import torch as t
|
||||
|
||||
import numpy as np
|
||||
import time
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
'''
|
||||
Implement ResidualBlock
|
||||
'''
|
||||
def __init__(self, inchannel, outchannel, stride=1, shortcut=None):
|
||||
super(ResidualBlock,self).__init__()
|
||||
self.left=nn.Sequential(
|
||||
nn.Conv2d(inchannel, outchannel, 3, stride, 1 , bias=False),
|
||||
nn.BatchNorm2d(outchannel),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(outchannel, outchannel, 3,1,1,bias=False),
|
||||
nn.BatchNorm2d(outchannel))
|
||||
self.right=shortcut
|
||||
|
||||
def forward(self,x):
|
||||
out=self.left(x)
|
||||
residual=x if self.right is None else self.right(x)
|
||||
out+=residual
|
||||
return F.relu(out)
|
||||
|
||||
class ResNet(nn.Module):
|
||||
'''
|
||||
Implement for main Module:ResNet34
|
||||
ResNet34
|
||||
'''
|
||||
def __init__(self, num_classes=1000):
|
||||
super(ResNet,self).__init__()
|
||||
|
||||
self.pre=nn.Sequential(
|
||||
nn.Conv2d(3,64,7,2,3,bias=False),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.MaxPool2d(3,2,1)
|
||||
)
|
||||
#repeated layer,consist of 3,4,6,3 residual blocks
|
||||
self.layer1=self._make_layer(64, 128 , 3)
|
||||
self.layer2=self._make_layer(128, 256 , 4, stride=2)
|
||||
self.layer3=self._make_layer(256, 512 , 6, stride=2)
|
||||
self.layer4=self._make_layer(512, 512 , 3, stride=2)
|
||||
|
||||
#classfication fc layer
|
||||
self.fc=nn.Linear(512,num_classes)
|
||||
|
||||
def _make_layer(self, inchannel, outchannel, block_num, stride=1):
|
||||
'''
|
||||
build layers
|
||||
'''
|
||||
shortcut=nn.Sequential(
|
||||
nn.Conv2d(inchannel, outchannel, 1, stride, bias=False),
|
||||
nn.BatchNorm2d(outchannel))
|
||||
|
||||
layers=[]
|
||||
layers.append(ResidualBlock(inchannel, outchannel, stride, shortcut))
|
||||
|
||||
for i in range(1, block_num):
|
||||
layers.append(ResidualBlock(outchannel, outchannel))
|
||||
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def forward(self,x):
|
||||
x=self.pre(x)
|
||||
|
||||
x=self.layer1(x)
|
||||
x=self.layer2(x)
|
||||
x=self.layer3(x)
|
||||
x=self.layer4(x)
|
||||
|
||||
x=F.avg_pool2d(x,7)
|
||||
x=x.view(x.size(0),-1)
|
||||
return self.fc(x)
|
||||
|
||||
if __name__ == '__main__':
|
||||
'''
|
||||
test function
|
||||
'''
|
||||
model=ResNet()
|
||||
#print(model)
|
||||
|
||||
input= t.autograd.Variable(t.randn(1,3,224,224))
|
||||
o=model(input)
|
||||
print(o)
|
||||
Reference in New Issue
Block a user