图像分类实战
读数据
def read_file(path):
for i in tqdm(range(11)):
file_dir = path + "/%02d"%i
file_list = os.listdir(file_dir) #列出文件夹下所有文件名字
xi = np.zeros((len(file_list), HW, HW, 3))
yi = np.zeros(len(file_list))
for j, img_name in enumerate(file_list):
img_path = os.path.join(file_dir, img_name)
img = Image.open(img_path)
img = img.resize((224, 224))
xi[j, ...] = img
yi[j] = i
pass
if i == 0:
X = xi
Y = yi
else :
X = np.concatenate((X, xi), axis=0)
Y = np.concatenate((Y, yi), axis=0)
print("读到了%d个训练数据"%len(Y))
return X, Y
1.因为训练集中的样例很多,读的很慢,所以使用tdqm可以显示循环的进度
2.numpy的zeros方法第一参数希望是shape所以多维的时候,第一个参数要带一个括号
3.使用enumerate在遍历可迭代对象时可以返回索引
4.文件地址连接的时候使用os.path.join函数,可以自动考虑斜杠
5.使用了concatenate函数让11个labeld文件数据连接在一起
6.因为文件的img是512512所以使用resize改为224224
train_transform = transforms.Compose(
[
transforms.ToPILImage(), # 224, 224,3 模型 ,但转化为3,224,224
transforms.RandomResizedCrop(224), #随即放大并裁切
transforms.RandomRotation(50), # 50度内随机旋转
transforms.ToTensor()
]
)
val_transform = transforms.Compose(
[
transforms.ToPILImage(), # 224, 224,3 模型 ,但转化为3,224,224
transforms.ToTensor()
]
)
为了充分学习训练集样本数据,需要对训练集进行数据扩增,但验证集不需要扩增,因为需要保证验证集每次提供的数据一致,扩增有随机性
class food_dataset(Dataset):
def __init__(self, path, mode):
self.x, self.y = read_file(path)
self.y = torch.LongTensor(self.y) #标签转为长整型,因为在分类任务时label都是整数表示类别,不像回归时是小数
if mode == "train":
self.transform = train_transform
else:
self.transform = val_transform
def __getitem__(self, item):
return self.transform(self.x[item]), self.y[item] #self.x[item]是ndarray矩阵224*224*3现在transform后是tensor 3*224*224
def __len__(self):
return len(self.y)
class myModel(nn.Module):
def __init__(self,num_class):
super(myModel,self).__init__()
# 3*224*224 -》 512*7*7 -》拉直-》全连接
self.conv1 = nn.Conv2d(3, 64, 3, 1, 1) #64*224*224
self.bn1 = nn.BatchNorm2d(64)
self.relu = nn.ReLU()
self.pool1 = nn.MaxPool2d(2) #64*112*112
self.layer1 = nn.Sequential(
nn.Conv2d(64, 128, 3, 1, 1), # 128*112*112
nn.BatchNorm2d(128),
nn.ReLU(),
nn.MaxPool2d(2) # 128*56*56
)
self.layer2 = nn.Sequential(
nn.Conv2d(128, 256, 3, 1, 1),
nn.BatchNorm2d(256),
nn.ReLU(),
nn.MaxPool2d(2) # 256*28*28
)
self.layer3 = nn.Sequential(
nn.Conv2d(256, 512, 3, 1, 1),
nn.BatchNorm2d(512),
nn.ReLU(),
nn.MaxPool2d(2) # 512*14*14
)
self.pool2 = nn.MaxPool2d(2) #512*7*7=25088
self.fc1 = nn.Linear(25088, 1000)
self.relu2 = nn.ReLU()
self.fc2 = nn.Linear(1000, num_class)
def forward(self,x):
x = self.conv1(x)
x = self.bn1(x)
x = self.relu(x)
x = self.pool1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.pool2(x)
x = x.view(x.size()[0], -1)
x = self.fc1(x)
x = self.relu2(x)
x = self.fc2(x)
return x
需要注意的是mymodel中forward在进行全连接前,需要使用view拉直一次
小批量随机梯度下降 (Mini-batch SGD) - 实际中最常用的:
这是 BGD 和 SGD 的折中方案。这也是 PyTorch 中 optim.SGD 默认优化的方式。
在每次迭代时,它会随机选择一个 mini-batch (一小批样本,例如 32, 64, 128 个)。
计算这个 mini-batch 内所有样本的平均损失,并求导,得到一个基于这个 mini-batch 的梯度估计。
更新公式: θ = θ - η * (1/|B|) * Σi∈B ∇J(θ; xi, yi) (其中 B 是当前的 mini-batch)
特点: 这个梯度仍然是一个估计值,带有一定的随机性(因为 mini-batch 是随机抽取的),但它比单个样本的梯度更稳定、更准确。同时,它又能利用现代硬件(GPU)的并行计算能力来加速计算 mini-batch 的梯度。这个 mini-batch 的梯度就是您在 PyTorch 中 loss.backward() 所计算出的那个梯度。

momentum采用了指数加权平均
Adam和AdamW

迁移学习使用大佬的模型和参数
from torchvision.models import resnet18
model = resnet18(pretrained=True) 使用大佬的参数
in_features = model.fc.in_features 替换大佬的分类头,resnet18只有最后的分类头是fc
model.fc = nn.Linear(in_features, 11)
李哥还在第五节课的代码文件里提前写好了了initialize_model方法
def initialize_model(model_name, num_classes, linear_prob=False, use_pretrained=True):
elif model_name == "resnet18":
""" Resnet18
"""
model_ft = models.resnet18(pretrained=use_pretrained) # 从网络下载模型 pretrain true 使用参数和架构, false 仅使用架构。
set_parameter_requires_grad(model_ft, linear_prob) # 是否为线性探测,线性探测: 固定特征提取器不训练。
其中的线性探测决定了模型的特征提取器的参数是否冻住
此外在第5节课的main文件中调用modle文件中的3个模块
其中model模块的transform提供了对文件照片的各种特征扩增,然后
if __name__ == '__main__': #运行的模块, 如果你运行的模块是当前模块
print("你运行的是data.py文件")
filepath = '../food-11_sample'
train_loader = getDataLoader(filepath, 'train', 8)
for i in range(3):
samplePlot(train_loader,True,isbat=False,ori=True)
val_loader = getDataLoader(filepath, 'val', 8)
for i in range(100):
samplePlot(val_loader,True,isbat=False,ori=True)

浙公网安备 33010602011771号