ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

将文件夹中的图片与CSV文件中的标签一一对应

将文件夹中的图片与CSV文件中的标签一一对应 import os import pandas as pd from torchvision import transforms from torch.utils.data import Dataset,DataLoader from PIL import Image import matplotlib.pyplot as plt import numpy as np def collect_data(image_folder,label_csv,output_csv):#命名一个函数 image_filesos.listdir(image_folder)#获得文件夹里每个文件的名字 print(image_files) dfpd.read_csv(label_csv)#读取CSV文件 image_data[]#这里统计每个文件与标签的对应情况 for image,label in zip(image_files,df[label]):#遍历每个图片与标签 image_data.append((os.path.join(image_folder, image),label)) print(image_data ) datapd.DataFrame(image_data,columns[file_path,label])#将他们记录在列表中并转变成矩阵 data.to_csv(output_csv)#储存在csv文件中 Ainput_images Blabel1.csv Coutput.csv collect_data(A,B,C)#应用函数得到输出 def calculate_mean_std(output_csv):#对输出CSV文件中记录的图片信息进行获取均值与标准差进行后续的图片转换 datapd.read_csv(output_csv) image_files data[file_path]#获取文件中的图片路径 pixel_values [] for image_file in image_files: with Image.open(image_file) as img: img img.convert(RGB) pixel_values.append(np.array(img)) pixel_values np.vstack(pixel_values) mean np.mean(pixel_values / 255.0, axis(0, 1, 2))#这里表示计算三个维度的均值与标准差 std np.std(pixel_values / 255.0, axis(0, 1, 2)) return mean, std mean,stdcalculate_mean_std(C)#为均值与标准差赋值 print(mean) class my_dataset(Dataset):#定义数据集的类 def __init__(self, csv_file, transformNone): self.datapd.read_csv(csv_file )#获得数据这里就是在CSV文件里得到图片与标签 self.transformtransform#定义图片转换的方法 def __len__(self):#获得数据的长度 return len(self.data) def __getitem__(self, idx):#得到每个标签与图片并返回 image_pathself.data.iloc[idx][file_path] labelself.data.iloc[idx][label] imageImage.open(image_path).convert(RGB) if self.transform: imageself.transform(image) return image , label transformtransforms.Compose ([transforms.Resize((224,224)),transforms.ToTensor(),transforms.Normalize(mean mean,std std)]) datasetmy_dataset(output.csv,transformtransform) dataloader DataLoader(dataset, batch_size1, shuffleTrue)#加载数据集 # 迭代 DataLoader 并显示图像和标签#下边我就不会了得到的图像是没有逆归一化的会黑乎乎的 for i, (image, label) in enumerate(dataloader): plt.imshow(image.squeeze().permute(1, 2, 0)) # 将张量转换为 numpy 数组并显示图像 plt.title(fLabel: {label.item()}) # 显示标签 plt.show() # 如果只想查看前几个样本可以添加一个条件退出循环 if i 4: # 例如只查看前5个样本 break
返回列表