news 2026/8/31 9:49:16

第五天 分类任务学习

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
第五天 分类任务学习

第五天 分类任务学习

今天要完成的是 视频分类任务,通过一系列图片数据(带标签,不带标签)来训练模型,这叫做 半监督学习。这次只做了数据预处理的工作,就是写两个dataset的类,怎么加载数据
数据集分为food_dataset和semi_dataset,一个训练带标签数据,一个不带标签数据,两种数据集

随机种子:

有点类似我的世界那种世界种子,有了种子,开局就是固定的地图
defseed_everything(seed):torch.manual_seed(seed)torch.cuda.manual_seed(seed)torch.cuda.manual_seed_all(seed)torch.backends.cudnn.benchmark=Falsetorch.backends.cudnn.deterministic=Truerandom.seed(seed)np.random.seed(seed)os.environ['PYTHONHASHSEED']=str(seed)#################################################################seed_everything(0)###############################################

这个 seed_everything 函数把"所有可能产生随机性的来源"都固定住了,目的是让训练过程完全可复现:
固定的随机源:

  1. torch.manual_seed - PyTorch CPU 上的随机操作(权重初始化、Dropout 等)
  2. torch.cuda.manual_seed(_all) - GPU 上的随机操作(单卡/多卡)
  3. random.seed - Python 标准库的随机函数(如 random.shuffle)
  4. np.random.seed - NumPy 的随机操作(数据预处理里可能用到)
  5. PYTHONHASHSEED - Python 字典哈希的随机性(影响某些遍历顺序)
    关闭的不确定性:
  6. cudnn.benchmark = False - 禁用自动选最快算法(不同算法结果可能微小差异)
  7. cudnn.deterministic = True - 强制用确定性算法(牺牲一点速度换完全一致的结果)

作用场景:

  1. 权重初始化每次一样
  2. DataLoader 的 shuffle 每次顺序一样
  3. Dropout 每次丢弃的神经元一样
  4. 数据增强(如果用了带随机的 transform)每次也能复现
train_transform=transforms.Compose([transforms.ToPILImage(),#224, 224, 3模型 :3, 224, 224transforms.RandomResizedCrop(224),transforms.RandomRotation(50),transforms.ToTensor()])val_transform=transforms.Compose([transforms.ToPILImage(),#224, 224, 3模型 :3, 224, 224transforms.ToTensor()])
  1. transforms.ToPILImage()是将张量转换为pil图片格式的函数
  2. transforms.RandomResizedCrop(224) 是 随机裁剪图像的一部分,并缩放到指定大小,提升模型鲁棒性
  3. transforms.RandomRotation()是用来旋转图片的
  4. 然后再次转换为张量

food_dataset

# 定义food_dataset类,支持三种模式,train, val, semiclassfood_dataset(Dataset):def__init__(self,path,mode="train"):self.mode=modeifmode=="semi":# 无标签模式:只读图片,不读标签self.X=self.read_files(path)else:# 有标签模式:读图片和标签self.X,self.Y=self.read_files(path)self.Y=torch.LongTensor(self.Y)# 标签转为长整型# 根据模式选择数据增强策略self.transform=train_transformifmode=="train"elseval_transformdefread_files(self,path):ifself.mode=="semi":# 无标签模式:直接读取文件夹下所有图片file_list=os.listdir(path)X=np.zeros((len(file_list),HW,HW,3))forj,img_nameinenumerate(file_list):img_path=os.path.join(path,img_name)img=Image.open(img_path)# 用PIL读图片img=img.resize((HW,HW))# 统一resize到224x224X[j,...]=np.array(img)print("读到了 %d 个无标签数据"%len(X))returnXelse:# 有标签模式:遍历11个子文件夹(00-10)foriintqdm(range(11)):file_dir=os.path.join(path,"%02d"%i)# 子文件夹名:00, 01, ..., 10file_list=os.listdir(file_dir)# 预分配数组存储当前类别的所有图片xi=np.zeros((len(file_list),HW,HW,3),dtype=np.uint8)yi=np.zeros(len(file_list),dtype=np.uint8)# 读取当前类别的所有图片forj,img_nameinenumerate(file_list):img_path=os.path.join(file_dir,img_name)img=Image.open(img_path)# 用PIL读图片img=img.resize((HW,HW))# 统一resize到224x224xi[j,...]=np.array(img)yi[j]=i# 文件夹序号就是类别标签# 拼接到总数组ifi==0:X=xi Y=yielse:X=np.concatenate((X,xi),axis=0)# 在第0维拼接上xi, xi的shape是(当前类别图片数, 224, 224, 3), 拼接后X的shape是(之前类别总图片数 + 当前类别图片数, 224, 224, 3)Y=np.concatenate((Y,yi),axis=0)# 在第0维拼接上yi, yi的shape是(当前类别图片数,), 拼接后Y的shape是(之前类别总图片数 + 当前类别图片数,)print("读到了 %d 个有标签数据"%len(Y))returnX,Ydef__getitem__(self,index):ifself.mode=="semi":# 无标签模式:返回 (变换后的图, 原始图)returnself.transform(self.X[index])else:# 有标签模式:返回 (变换后的图, 标签)returnself.transform(self.X[index]),self.Y[index]def__len__(self):returnlen(self.X)

BUG

  1. PermissionError: [Errno 13] Permission denied: ‘C:\baidunetdiskdownload\ai资料\李哥深度学习\04.课程代码\第四五节_分类代码\food_classification\food-11\training\labeled’
    open(path, “r”) 是用来打开文件的,但是这里打开了文件夹
  2. 不能用open(path, “r”)读取图片
fromPILimportImage Image.open(path)
  1. Conda在下载安装包时报错:
    PackagesNotFoundError: The following packages are not available from current channels:
    XXXXXX(包名)
    有如下两种解决方法:
    方法一:将conda-forge添加到搜索路径上
    在命令行运行下方指令,然后重新安装。
    conda config --append channels conda-forge
    conda install 需要安装的包名
# 定义semi_dataset类,专门用于无标签数据的读取, 并且使用模型预测标签,返回 (变换后的图, 预测标签) 的形式classsemi_dataset(Dataset):def__init__(self,no_label_loder,model,device,thres=0.99):x,y=self.get_label(no_label_loder,model,device,thres)# 获取预测标签,thres是置信度阈值,可以根据需要调整ifx==[]:self.flag=False# 没有满足条件的样本print("没有满足条件的样本")else:self.flag=True# 有满足条件的样本self.X=np.array(x)# 满足条件的图像数据self.Y=torch.LongTensor(y)# 满足条件的预测标签转为长整型self.transform=train_transform# 半监督数据也使用训练集的变换策略defread_files(self,path):file_list=os.listdir(path)X=np.zeros((len(file_list),HW,HW,3))forj,img_nameinenumerate(file_list):img_path=os.path.join(path,img_name)img=Image.open(img_path)# 用PIL读图片img=img.resize((HW,HW))# 统一resize到224x224X[j,...]=np.array(img)print("读到了 %d 个无标签数据"%len(X))returnXdef__getitem__(self,index):# 无标签模式:返回 (变换后的图, 预测标签)returnself.transform(self.X[index]),self.Y[index]def__len__(self):returnlen(self.X)defget_label(self,no_label_loder,model,device,thres):model=model.to(device)pred_prob=[]labels=[]x=[]y=[]soft=nn.Softmax()withtorch.no_grad():forbat_x,_inno_label_loder:bat_x=bat_x.to(device)pred=model(bat_x)prob=soft(pred)# 计算概率prob_max,pred_value=prob.max(1)pred_prob.extend(prob_max.cpu().numpy())# 将概率添加到列表中labels.extend(pred_value.cpu().numpy())foriinrange(len(pred_prob)):ifpred_prob[i]>=thres:# 如果概率大于等于阈值x.append(no_label_loder.dataset.X[i])# 添加对应的图像数据y.append(labels[i])# 添加对应的预测标签returnx,y

通过read_files()读取数据存储在x中,在get_label()中完成预测,然后根据置信度筛选出99%以上的样例作为x,y

Adam和AdamW区别

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/7/14 17:20:29

Home Assistant Python自定义集成开发:从入门到精通——Config Flow、Entity、Service与事件总线深度解析

Home Assistant Python自定义集成开发:从入门到精通——Config Flow、Entity、Service与事件总线深度解析 Home Assistant Python自定义集成开发:从入门到精通——Config Flow、Entity、Service与事件总线深度解析 引言:智能家居的“最后一公里” 第一章:架构概览与环境准备…

作者头像 李华
网站建设 2026/7/14 17:20:28

Claude Code 实战命令

#ClaudeCode实战指南 #高效编程助手 Token管理场景• 关键命令: • /clear:完全清空对话历史(适合任务切换) • /compact:智能压缩历史记录(保留关键摘要) • /context:实时显示Toke…

作者头像 李华
网站建设 2026/7/14 17:20:29

Docker 容器

容器是 Docker核心概念。 简单的说,容器是独立运行的一个或一组应用,以及它们的运行环境。 对应的,虚拟机可以理解为模拟运行的一整套操作系统(提供了运行态环境和其他系统环境)和运行在上面的应用。咱们可以从对象和类的关系来看容器与镜像的…

作者头像 李华
网站建设 2026/7/14 17:20:32

GESP C++考试五级语法知识(二、埃氏筛和线性筛)

🌟《素数王国的两种超级筛子》故事前言:1、在数字王国里,国王需要一份 素数名单。(1)比如:1 ~ 30 之间的所有素数(2)答案应该是:2 3 5 7 11 13 17 19 23 29(3…

作者头像 李华
网站建设 2026/7/14 17:20:32

序:Hello, Robot!

具身智能(Embodied AI),是2026年最硬核、最性感、但也最容易让人一头雾水的赛道。每天都有新的论文砸向 Arxiv,硅谷的初创公司融资新闻刷屏,但很少有人能给你一套从硬件到算法、从传统到SOTA的、系统化的技术全栈梳理。…

作者头像 李华