{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30635,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-01T06:12:22.754971Z","iopub.execute_input":"2024-03-01T06:12:22.755407Z","iopub.status.idle":"2024-03-01T06:12:22.761935Z","shell.execute_reply.started":"2024-03-01T06:12:22.755375Z","shell.execute_reply":"2024-03-01T06:12:22.760563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\n\nbasic_root = '/kaggle/input/cassava-leaf-disease-classification'\n#在上面的大文件中，train_images里是jpg的本身，而train.csv则展示了每张照片的名称和对应的标签\n#先对train_images文件中的海量图片进行分析\n#遍历文件夹\nimage_list = os.listdir(os.path.join(basic_root,'train_images'))\nprint('total number picture is ',len(image_list))\n\n#每张图片的大小很重要，创建一个字典，返回图片的shape总共有多少类，每种类有多少张，这里为了避免卡顿，只查看了前300张\n#提示 dictionary.get(a,b):查找a,若不存在，返回b\nimage_shape = {}\nfor image_name in image_list[0:300]:\n    image = cv2.imread(os.path.join(basic_root,'train_images',image_name))\n    image_shape[image.shape] = image_shape.get(image.shape,0)+1\nprint(image_shape)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:24.351366Z","iopub.execute_input":"2024-03-01T06:12:24.351761Z","iopub.status.idle":"2024-03-01T06:12:27.427461Z","shell.execute_reply.started":"2024-03-01T06:12:24.351731Z","shell.execute_reply":"2024-03-01T06:12:27.426308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#接下来读取train.csv文件，把其变成dataframe的格式，应用pandas库中的不同方法，可以实现对数据的快速的清洗\nimage_csv = pd.read_csv(os.path.join(basic_root,'train.csv'))\nprint(image_csv)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:27.429675Z","iopub.execute_input":"2024-03-01T06:12:27.430144Z","iopub.status.idle":"2024-03-01T06:12:27.461974Z","shell.execute_reply.started":"2024-03-01T06:12:27.430101Z","shell.execute_reply":"2024-03-01T06:12:27.460701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#label的代号与label的真实名字的对应，已经写在了json文件中，我们可以让csv文件变得更加充实一点\nimport json\nimage_json = json.load(open(os.path.join(basic_root,'label_num_to_disease_map.json')))\nprint('image_json is:',image_json)\n#我们的想法是利用map函数，把pdframe新增一列出来，用来存放数字label所对应的实际含义\n#1.把json变成字典，记得将键的类型变成int型\nimage_json_tolabel = {int(i):j for i,j in image_json.items()}\nprint('image_json_tolabel is',image_json_tolabel)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:29.566739Z","iopub.execute_input":"2024-03-01T06:12:29.567138Z","iopub.status.idle":"2024-03-01T06:12:29.581393Z","shell.execute_reply.started":"2024-03-01T06:12:29.567107Z","shell.execute_reply":"2024-03-01T06:12:29.580150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#启动\nimage_csv['label_name'] =image_csv['label'].map(image_json_tolabel)\nprint(image_csv)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:31.537011Z","iopub.execute_input":"2024-03-01T06:12:31.537429Z","iopub.status.idle":"2024-03-01T06:12:31.557733Z","shell.execute_reply.started":"2024-03-01T06:12:31.537397Z","shell.execute_reply":"2024-03-01T06:12:31.556645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#csv变得顺眼多了以后，我们要看一看，里面总共有多少类，每一个类的图片的比例是多少\nimport seaborn as sn\nimport matplotlib.pyplot as plt\nplt.figure(figsize = (6,6))\nsn.countplot(y = 'label_name',data = image_csv)\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:33.323515Z","iopub.execute_input":"2024-03-01T06:12:33.323969Z","iopub.status.idle":"2024-03-01T06:12:34.326826Z","shell.execute_reply.started":"2024-03-01T06:12:33.323934Z","shell.execute_reply":"2024-03-01T06:12:34.325645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#还可以有一些更为精细的可视化操作\nimport math\ndef visualize_batch(image_names,labels):\n    plt.figure(figsize=(18,18))\n    #有顺序地画出image_name对应的图片，将label作为图片的title\n    for index,(image_name,label) in enumerate(zip(image_names,labels)):\n        plt.subplot(int(math.sqrt(len(image_names))),int(math.sqrt(len(image_names))), index + 1)\n        image = cv2.imread(os.path.join(basic_root,'train_images',image_name))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        print(image.shape)\n        plt.title(f\"Class: {label}\", fontsize=12)\n        plt.imshow(image)\n\n    plt.show()\n#从原始csv中sample9个样本   \ntmp_csv = image_csv.sample(9)\n#注意，visualize_batch输入的参数，类型是list,dataframe操作中非常多种方法可以把某一行或者某一列变成list,或者array,\n#这也是我个人觉得在处理数据中pd方便的原因\nvisualize_batch(tmp_csv['image_id'].values,tmp_csv['label_name'].values)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:36.371715Z","iopub.execute_input":"2024-03-01T06:12:36.372148Z","iopub.status.idle":"2024-03-01T06:12:39.851701Z","shell.execute_reply.started":"2024-03-01T06:12:36.372113Z","shell.execute_reply":"2024-03-01T06:12:39.847869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#下面开始构造dataset和dataloader\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\n#把 csv中，’image_id'的一列的10个元素，以及‘label'的一列的10个元素，单独提取出来，组成一个新list,\n#其实这些操作完全可以在dataset init方法中完成，只是为了让dataset更加简单和方便理解\ndata_image = image_csv['image_id'].values\ntarget = image_csv['label'].values\ndata = list(zip(data_image[:10],target[:10]))\n\nimport cv2\nfrom PIL import Image\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:45.926656Z","iopub.execute_input":"2024-03-01T06:12:45.927346Z","iopub.status.idle":"2024-03-01T06:12:45.933678Z","shell.execute_reply.started":"2024-03-01T06:12:45.927310Z","shell.execute_reply":"2024-03-01T06:12:45.932674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#我刚刚接触深度学习的时，一直对dataset有种恐惧感，不知道它到底是什么，怎么使用，现在模模糊糊的对它有了一点感觉，dataset是一个\n#框架，对于一般的任务，只需要覆写Dataset的3个方法，就可以让输入的数据，按照你想要的方式，整整齐齐的走出来。尤其是__getitem__方法\n#它决定了，你的数据是如何被喂给模型的。如果非让我用一个生活中的例子来表述的话，我会把它比成切猪肉，鉴于大家都喜欢买瘦肉，\n#所以说，一刀下去，大范围的瘦肉（data）和小范围的肥肉（label）都被割了一块下来。\n\n#当写到这里的时候，我产生了一个疑惑，PIL和cv2读取图片到底有什么区别。我找到了如下两个例子。\nfrom torchvision.transforms import transforms\nclass vision_dataset(Dataset):\n    def __init__(self,data,use_cv2,transform = None):\n        self.data = data\n        self.transform = transforms.Compose([\n            transforms.ToTensor()      # 这里仅以最基本的为例\n        ])\n        self.use_cv2 = use_cv2\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self,index):\n        image = self.data[index][0]\n        if self.use_cv2:\n            image = cv2.imread(os.path.join(basic_root,'train_images',image))#读取的是BGR数据\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)#转成RGB模式\n            #以上两个的顺序都是H,W,C。需要转化为C,H,W\n            image = torch.from_numpy(image).permute(2, 0, 1)/255\n        else:\n            image = Image.open(os.path.join(basic_root,'train_images',image))  # 读取到的是RGB， W, H, C\n            image = self.transform(image)   # transform转化image为：C, H, W\n\n        label = self.data[index][1]\n                \n        return image,label\n    \n#实例化\npil_dataset = vision_dataset(data,use_cv2 = False)\ncv2_dataset = vision_dataset(data,use_cv2 = True)\n\n#取出两个不同的dataset序列的第一个组合。\nimage_cv2,label_cv2 = cv2_dataset[0]\nimage_pil,label_pil = pil_dataset[0]\n\nprint(image_pil)\nprint(image_cv2)\nprint(torch.sum(image_cv2-image_pil,dim = (0,1,2)))\n#感兴趣的可以看一看。读出来的据说还是略微有一些区别的（虽然我读出来的没啥区别），主要是库的原因。\n#为了和其他人保持一致，也为了写起来方便。最重要的是传言PIL读出来的图片\n#训练容易收敛。建议使用PIL方法。\n#注意在实际操作中，dataset还缺少一部分，就是在test_epoch的时候应该怎么写，在之后的test的构建中，dataset __getitem__不能返回label了，因为test_data一定是没有label的，\n#不过无伤大雅，这个notebook的问题不是在于讨论这个，之后我会出一个比较完整的训练流程。\n\n\n\n    \n        \n        \n        ","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:48.411048Z","iopub.execute_input":"2024-03-01T06:12:48.411482Z","iopub.status.idle":"2024-03-01T06:12:48.470337Z","shell.execute_reply.started":"2024-03-01T06:12:48.411448Z","shell.execute_reply":"2024-03-01T06:12:48.469234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#写完dataset,dataloader就是一行代码的事情。还是拿猪肉做比喻，这个loader相当于是在选择切多厚的肉片（batcg_size）。\ndataloader = DataLoader(pil_dataset, batch_size=2, shuffle=False, num_workers=1)\n\n#在我看来，dataloader就是一个小型的dataset。从dataloader里面取数据，我见过的一般有两种,以取出label为例\n#1.\nfor i, (images, labels) in enumerate(dataloader):\n    print(labels)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:52.624377Z","iopub.execute_input":"2024-03-01T06:12:52.624765Z","iopub.status.idle":"2024-03-01T06:12:53.008575Z","shell.execute_reply.started":"2024-03-01T06:12:52.624736Z","shell.execute_reply":"2024-03-01T06:12:53.007128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#2.\niter_dataloader = iter(dataloader)\nfor i in range(int(len(pil_dataset)/2)):\n    images, labels = next(iter_dataloader)\n    print(labels)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:54.822148Z","iopub.execute_input":"2024-03-01T06:12:54.822839Z","iopub.status.idle":"2024-03-01T06:12:55.051840Z","shell.execute_reply.started":"2024-03-01T06:12:54.822793Z","shell.execute_reply":"2024-03-01T06:12:55.050060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#如果你记性不好的话，咱们可以把dataset中的label打印出来\nfor i in range(len(pil_dataset)):\n    images,labels = pil_dataset[i]\n    print(labels)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:12:56.652992Z","iopub.execute_input":"2024-03-01T06:12:56.653503Z","iopub.status.idle":"2024-03-01T06:12:56.755632Z","shell.execute_reply.started":"2024-03-01T06:12:56.653462Z","shell.execute_reply":"2024-03-01T06:12:56.754250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#可以很清楚的看到，dataloader的label就是dataset的label以batch_size的大小成批输出的。类型略微有区别，dataloader输出的是\n#tensor型的变量。","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#好困。。。。。下一篇我想讨论一个如何更加自由的定义一个dataset,今天就到这里了。","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"OK了家人们，早上起来神清气爽，让我们继续探究如何更加自由的定义dataset，那就到了我们今天的主题，自定义sampler。","metadata":{}},{"cell_type":"markdown","source":"sampler顾名思义，就是一个采样器，决定着dataloader在batch_size固定的情况下，取哪几个数据，在简单的情况下，sampler都不需要自己定义，因为pytoch自己本身就给我们提供了两种sampler,一种是顺序采样器，一种就是随机采样器。随着dataloader中shuffle的参数变化而变化。","metadata":{}},{"cell_type":"code","source":"default_sampler = dataloader.sampler\nfor i in default_sampler:\n    print(i)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:00.588331Z","iopub.execute_input":"2024-03-01T06:13:00.588806Z","iopub.status.idle":"2024-03-01T06:13:00.595074Z","shell.execute_reply.started":"2024-03-01T06:13:00.588769Z","shell.execute_reply":"2024-03-01T06:13:00.593709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"如果你记性足够好的话，你就会记得，我们的dataset中就是包含了10个图片和其对应的label,而且我们实例化dataloader的时候，shuffle是为False的。顺序采样器就是按照内部的序号，依次取出dataset的数据。","metadata":{}},{"cell_type":"markdown","source":"让我们再来看看随机采样器。","metadata":{}},{"cell_type":"code","source":"dataloader = DataLoader(pil_dataset, batch_size=2, shuffle=True, num_workers=1)\ndefault_sampler = dataloader.sampler\nfor i in default_sampler:\n    print(i)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:02.756516Z","iopub.execute_input":"2024-03-01T06:13:02.756907Z","iopub.status.idle":"2024-03-01T06:13:02.765344Z","shell.execute_reply.started":"2024-03-01T06:13:02.756879Z","shell.execute_reply":"2024-03-01T06:13:02.764123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"所以说，shuffle这个参数，它与其说是把dataset里的数据打乱顺序，倒不如说是将数据的序号打乱，然后变成新队列，然后输出。","metadata":{}},{"cell_type":"markdown","source":"问题在这里出现，如果我不想用顺序，也不想用乱序，我想用我定义的规律来去实现数据的输出呢，比如我非常拧巴，我一定要让后一半的数据先进来，前一半的数据后进来呢？","metadata":{}},{"cell_type":"code","source":"import random\nfrom torch.utils.data.sampler import Sampler\n\nclass mysampler(Sampler):\n    def __init__(self,dataset):\n        halfway_point = int(len(dataset)/2)\n        self.first = list(range(halfway_point))\n        self.second = list(range(halfway_point,len(dataset)))\n\n    def __iter__(self):\n        random.shuffle(self.first)\n        random.shuffle(self.second)\n        return iter(self.second+self.first)\n\n    def __len__(self):\n        return len(dataset)\n\noursampler = mysampler(pil_dataset)\nfor i in oursampler:\n    print(i)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:05.374735Z","iopub.execute_input":"2024-03-01T06:13:05.375809Z","iopub.status.idle":"2024-03-01T06:13:05.385031Z","shell.execute_reply.started":"2024-03-01T06:13:05.375771Z","shell.execute_reply":"2024-03-01T06:13:05.384144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#把sampler放在dataloader中\ndataloader_shuflle_half = DataLoader(pil_dataset,sampler=oursampler,batch_size=3)\nfor i,(images,labels) in enumerate(dataloader_shuflle_half):\n    print(labels)\n#这里label没有提前控制好，每个图片的label设置的不一样的话可能会更加直观。其实看上一段代码我觉得已经说的很清楚了。","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:07.441401Z","iopub.execute_input":"2024-03-01T06:13:07.442115Z","iopub.status.idle":"2024-03-01T06:13:07.555953Z","shell.execute_reply.started":"2024-03-01T06:13:07.442081Z","shell.execute_reply":"2024-03-01T06:13:07.554793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"这里有一个小bug,我是把后半段打乱，前半段打乱，加到一起，作为输出，5个后半段序号，5个前半段序号，如果输出的时候，batch_size为2时，在第三个batch中，一定会有一个后半段的序号和前半段的序号一块输入进来。所以我们需要更精细的sampler,那就是batch_sampler。","metadata":{}},{"cell_type":"code","source":"\ndef chunk(indices,chunk_size):\n    return torch.split(torch.tensor(indices),chunk_size)\n\nclass batch_samplers(Sampler):\n    def __init__(self, dataset, batch_size):\n        half_way_point = int(len(dataset)/2)\n        self.first = list(range(half_way_point))\n        self.second = list(range(half_way_point,len(dataset)))\n        self.batch_size = batch_size\n\n    def __iter__(self):\n        random.shuffle(self.first)\n        random.shuffle(self.second)\n        first_batch = chunk(self.first,self.batch_size)\n        second_batch = chunk(self.second,self.batch_size)\n        combined = list(first_batch+second_batch)\n        combined = list(batch.tolist() for batch in combined)\n        #random.shuffle(combined)\n        return iter(combined)\n\n    def __len__(self):\n        return len(dataset)/batch_size\n\nbatch_size = 2\nmy_batch_sampler = batch_samplers(pil_dataset,batch_size)\nfor x in my_batch_sampler:\n    print(x)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:10.128191Z","iopub.execute_input":"2024-03-01T06:13:10.128662Z","iopub.status.idle":"2024-03-01T06:13:10.141295Z","shell.execute_reply.started":"2024-03-01T06:13:10.128629Z","shell.execute_reply":"2024-03-01T06:13:10.139999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"这样操作，以后每个batch里面只有前半段的序号，或者后半段的序号，两者不会一同出现。","metadata":{}},{"cell_type":"code","source":"dataloader_shuflle_half = DataLoader(pil_dataset,batch_sampler=my_batch_sampler)\nfor i,(images,labels) in enumerate(dataloader_shuflle_half):\n    print(labels)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:13.958240Z","iopub.execute_input":"2024-03-01T06:13:13.958719Z","iopub.status.idle":"2024-03-01T06:13:14.070858Z","shell.execute_reply.started":"2024-03-01T06:13:13.958683Z","shell.execute_reply":"2024-03-01T06:13:14.069468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"下面自定义一个比较有意义的sampler,在有一些深度学习方法中，比如通过生成新样本来平衡数据集的时候，在batch层面操作的时候，一般对batch的数据内容是有要求的。要么提前把数据整理的特别符合要求。要么就用自定义sampler了。\n\n已经该数据集的label总共有5类，如果我想要每一个batch的size为5，且每一个batch的数据刚好包含了这5种，一种一个(有fewshot那味了）。那应该怎么做呢？","metadata":{}},{"cell_type":"code","source":"from collections import OrderedDict\nimport random\n#把csv的所有的数据全部变成data\ndata_image = image_csv['image_id'].values\ntarget = image_csv['label'].values\ndata = list(zip(data_image,target))\nclass vision_dataset(Dataset):\n    def __init__(self,data,use_cv2,transform = None):\n        self.data = data\n        self.transform = transforms.Compose([\n            transforms.ToTensor()      # 这里仅以最基本的为例\n        ])\n        self.use_cv2 = use_cv2\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self,index):\n        image = self.data[index][0]\n        if self.use_cv2:\n            image = cv2.imread(os.path.join(basic_root,'train_images',image))#读取的是BGR数据\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)#转成RGB模式\n            #以上两个的顺序都是H,W,C。需要转化为C,H,W\n            image = torch.from_numpy(image).permute(2, 0, 1)/255\n        else:\n            image = Image.open(os.path.join(basic_root,'train_images',image))  # 读取到的是RGB， W, H, C\n            image = self.transform(image)   # transform转化image为：C, H, W\n\n        label = self.data[index][1]\n                \n        return image,label\n#实例化\npil_dataset = vision_dataset(data,use_cv2 = False)\n\nclass N_Way_K_Shot_BatchSampler(Sampler):\n    def __init__(self, data, max_iter):\n        _,self.y = zip(*data)\n        self.y = list(self.y)\n        self.max_iter = max_iter\n        self.label_dict = self.build_label_dict()\n        self.unique_classes_from_y = list(set(self.y))\n    #构建一个字典，key值为label，value值为拥有相同label的样本在dataset中的序号的集合\n    def build_label_dict(self):\n        label_dict = OrderedDict()\n        for i, label in enumerate(self.y):\n            if label not in label_dict:\n                label_dict[label] = [i]\n            else:\n                label_dict[label].append(i)\n        #print(label_dict)\n        return label_dict\n    #从字典中随机选取一个key值为cls的value(即序号)\n    def sample_examples_by_class(self, cls):\n        if cls not in self.unique_classes_from_y:\n            return []\n        sampled_examples = random.sample(self.label_dict[cls],1)  # sample without replacement\n      \n        return sampled_examples\n    #构造一个迭代器，通过调用next()方法不断生成符合条件的序号\n    def __iter__(self):\n        for _ in range(self.max_iter):\n            batch = []\n            classes = self.unique_classes_from_y\n            for cls in classes:\n                samples_for_this_class = self.sample_examples_by_class(cls)\n                batch.extend(samples_for_this_class)\n            yield batch\n\n    def __len__(self):\n        return self.max_iter\n    \n    \n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:16.437400Z","iopub.execute_input":"2024-03-01T06:13:16.437831Z","iopub.status.idle":"2024-03-01T06:13:16.468725Z","shell.execute_reply.started":"2024-03-01T06:13:16.437794Z","shell.execute_reply":"2024-03-01T06:13:16.467372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_way_k_shot_sampler = N_Way_K_Shot_BatchSampler(data,10000)\n# for i in n_way_k_shot_sampler:\n#     print(i)\n\ndataloader_complex = DataLoader(pil_dataset,batch_sampler=n_way_k_shot_sampler)\niter_dataloader = iter(dataloader_complex)\n\n\n#迭代两次\nfor i in range(2):\n    images,labels = next(iter_dataloader)\n    #把tensor（【5，3，600，800】）重新转化为图片\n    #改变维度顺序\n    batch_images = images.permute(0, 2, 3, 1)\n\n    # 设置子图布局\n    fig, axes = plt.subplots(1, 5, figsize=(15, 3))\n\n    # 循环显示每张图片\n    for j in range(5):\n        axes[j].imshow(batch_images[j].numpy())\n        axes[j].axis('off')  # 可选：关闭坐标轴\n        axes[j].set_title(str(labels[j].item()))\n\n    plt.show()\n#     plt.close()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:19.730477Z","iopub.execute_input":"2024-03-01T06:13:19.730925Z","iopub.status.idle":"2024-03-01T06:13:22.048410Z","shell.execute_reply.started":"2024-03-01T06:13:19.730891Z","shell.execute_reply":"2024-03-01T06:13:22.047359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"结束，收工。","metadata":{}},{"cell_type":"code","source":"data_image = image_csv['image_id'].values\ntarget = [[1,2],[1],[3,4,5],[6]]\n#target = image_csv['label'].values\ndata = list(zip(data_image[:4],target))\nprint(data)","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:13:25.538773Z","iopub.execute_input":"2024-03-01T06:13:25.539243Z","iopub.status.idle":"2024-03-01T06:13:25.546501Z","shell.execute_reply.started":"2024-03-01T06:13:25.539206Z","shell.execute_reply":"2024-03-01T06:13:25.545465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_,y = zip(*data)\nprint(y)\nlist_y = list(y)\nprint(list(y))\ndef label_unique(labels):\n    list_unique = []\n    for i in y:\n        for j in i:\n            list_unique.append(j)\n    return list(set(list_unique))\n\nprint(label_unique(list_y))\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:04:08.337122Z","iopub.execute_input":"2024-03-01T06:04:08.338222Z","iopub.status.idle":"2024-03-01T06:04:08.345717Z","shell.execute_reply.started":"2024-03-01T06:04:08.338181Z","shell.execute_reply":"2024-03-01T06:04:08.344330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import OrderedDict\n# _,y = zip(*data)\n# #print(list(y))\n# list_y = list(y)\n# label_dict = OrderedDict()\n# for i, label in enumerate(list_y):\n#     for j in label:\n#         if j not in label_dict:\n#             label_dict[j] = [i]\n#         else:\n#             label_dict[j].append(i)\n# print(label_dict)\n#OrderedDict([(1, [0, 1]), (2, [0]), (3, [2]), (4, [2]), (5, [2]), (6, [3])])\n#需要注意的是，有一些域的label数为0的\nclass N_Way_K_Shot_BatchSampler(Sampler):\n    def __init__(self, data, max_iter):\n        _,self.y = zip(*data)\n        self.y = list(self.y)\n        self.max_iter = max_iter\n        self.label_dict = self.build_label_dict()\n        self.unique_classes = list(range(1,8))\n        self.unique_classes_from_y = label_unique(self.y)\n    #构建一个字典，key值为label，value值为拥有相同label的样本在dataset中的序号的集合\n    def build_label_dict(self):\n        label_dict = OrderedDict()\n        for i, label in enumerate(self.y):\n            for j in label:\n                if j not in label_dict:\n                    label_dict[j] = [i]\n                else:\n                    label_dict[j].append(i)\n        #print(label_dict)\n        return label_dict\n    #从字典中随机选取一个key值为cls的value(即序号)\n    def sample_examples_by_class(self, cls):\n        if cls not in self.unique_classes_from_y:\n            sampled_examples = [None]\n        else:\n            sampled_examples = random.sample(self.label_dict[cls],1)  # sample without replacement\n      \n        return sampled_examples\n    #构造一个迭代器，通过调用next()方法不断生成符合条件的序号\n    def __iter__(self):\n        for _ in range(self.max_iter):\n            batch = []\n            classes = self.unique_classes\n            for cls in classes:\n                samples_for_this_class = self.sample_examples_by_class(cls)\n                batch.extend(samples_for_this_class)\n            yield batch\n\n    def __len__(self):\n        return self.max_iter\n\nsampler_ar = N_Way_K_Shot_BatchSampler(data,10)\nfor i in sampler_ar:\n    print(i)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:46:18.656520Z","iopub.execute_input":"2024-03-01T06:46:18.656930Z","iopub.status.idle":"2024-03-01T06:46:18.674214Z","shell.execute_reply.started":"2024-03-01T06:46:18.656898Z","shell.execute_reply":"2024-03-01T06:46:18.672718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import OrderedDict\nimport random\n#把csv的所有的数据全部变成data\ndata_image = image_csv['image_id'].values\ntarget = [[1,2],[1],[3,4,5],[6]]\n#target = image_csv['label'].values\ndata = list(zip(data_image[:4],target))\n#data = list(zip(data_image,target))\nclass vision_dataset(Dataset):\n    def __init__(self,data,use_cv2,transform = None):\n        self.data = data\n        self.transform = transforms.Compose([\n            transforms.ToTensor()      # 这里仅以最基本的为例\n        ])\n        self.use_cv2 = use_cv2\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self,index):\n        image = self.data[index][0]\n        if self.use_cv2:\n            image = cv2.imread(os.path.join(basic_root,'train_images',image))#读取的是BGR数据\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)#转成RGB模式\n            #以上两个的顺序都是H,W,C。需要转化为C,H,W\n            image = torch.from_numpy(image).permute(2, 0, 1)/255\n        else:\n            image = Image.open(os.path.join(basic_root,'train_images',image))  # 读取到的是RGB， W, H, C\n            image = self.transform(image)   # transform转化image为：C, H, W\n\n        label = self.data[index][1]\n                \n        return image,label\n    \npil_dataset = vision_dataset(data,use_cv2 = False)    \nn_way_k_shot_sampler = N_Way_K_Shot_BatchSampler(data,10000)\n# for i in n_way_k_shot_sampler:\n#     print(i)\n\ndataloader_complex = DataLoader(pil_dataset,batch_sampler=n_way_k_shot_sampler)\n\niter_dataloader = iter(dataloader_complex)\n\n\nfor i in range(2):\n    images,labels = next(iter_dataloader)\n    #把tensor（【5，3，600，800】）重新转化为图片\n    #改变维度顺序\n    print(labels)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-01T06:17:31.847472Z","iopub.execute_input":"2024-03-01T06:17:31.847896Z","iopub.status.idle":"2024-03-01T06:17:32.135296Z","shell.execute_reply.started":"2024-03-01T06:17:31.847865Z","shell.execute_reply":"2024-03-01T06:17:32.133543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import OrderedDict\nimport random\n#把csv的所有的数据全部变成data\ndata_image = image_csv['image_id'].values\ntarget = [[1,2],[1],[3,4,5],[6]]\n#target = image_csv['label'].values\ndata = list(zip(data_image[:4],target))\n#data = list(zip(data_image,target))\nclass vision_dataset(Dataset):\n    def __init__(self,data,use_cv2,transform = None):\n        self.data = data\n        self.transform = transforms.Compose([\n            transforms.ToTensor()      # 这里仅以最基本的为例\n        ])\n        self.use_cv2 = use_cv2\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self,index):\n        image = self.data[index][0]\n        if self.use_cv2:\n            image = cv2.imread(os.path.join(basic_root,'train_images',image))#读取的是BGR数据\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)#转成RGB模式\n            #以上两个的顺序都是H,W,C。需要转化为C,H,W\n            image = torch.from_numpy(image).permute(2, 0, 1)/255\n        else:\n            image = Image.open(os.path.join(basic_root,'train_images',image))  # 读取到的是RGB， W, H, C\n            image = self.transform(image)   # transform转化image为：C, H, W\n\n        label = self.data[index][1]\n                \n        return image,label\n#实例化\npil_dataset = vision_dataset(data,use_cv2 = False)\n\nclass N_Way_K_Shot_BatchSampler(Sampler):\n    def __init__(self, data, max_iter):\n        _,self.y = zip(*data)\n        self.y = list(self.y)\n        self.max_iter = max_iter\n        self.label_dict = self.build_label_dict()\n        self.unique_classes_from_y = list(set(self.y))\n    #构建一个字典，key值为label，value值为拥有相同label的样本在dataset中的序号的集合\n    def build_label_dict(self):\n        label_dict = OrderedDict()\n        for i, label in enumerate(self.y):\n            for j in label\n                if j not in label_dict:\n                    label_dict[j] = [i]\n                else:\n                    label_dict[j].append(i)\n        #print(label_dict)\n        return label_dict\n    #从字典中随机选取一个key值为cls的value(即序号)\n    def sample_examples_by_class(self, cls):\n        if cls not in self.unique_classes_from_y:\n            return []\n        sampled_examples = random.sample(self.label_dict[cls],1)  # sample without replacement\n      \n        return sampled_examples\n    #构造一个迭代器，通过调用next()方法不断生成符合条件的序号\n    def __iter__(self):\n        for _ in range(self.max_iter):\n            batch = []\n            classes = self.unique_classes_from_y\n            for cls in classes:\n                samples_for_this_class = self.sample_examples_by_class(cls)\n                batch.extend(samples_for_this_class)\n            yield batch\n\n    def __len__(self):\n        return self.max_iter\n    \n    \n","metadata":{},"execution_count":null,"outputs":[]}]}