Files
modelscope-scepter/docs/zh_cn/scepter/dataset.md
T

2.4 KiB
Raw Blame History

数据集模块 (Dataset)

总览

在继承BaseDataset基础上注册每个task各自的数据集读取模块,BaseDataset中封装有File System以及transform pipeline;


scepter.modules.data.dataset.BaseDataset

function __init__()

Parameters(输入参数)

(cfg: scepter.modules.utils.config.Config, logger = None) -> None

config:

  • MODE —— [train, test, eval];
  • TRANSFORMS —— 用于构造pipeline处理每个样本,详见docs.transforms;
  • FILE_SYSTEM —— File IO Handler,支持不同类型文件的读写操作,详见docs.utils.file_clients;

object(内部)

  • self.pipeline —— transform pipeline
  • self.fs_prefix —— The prefix of instantiated fs_client
  • self.local_we —— local rank info when is_distributed

function __getitem__()

Parameters(输入参数):index

index的类别由sampler确定,详见data/registry.py;

一般默认的torch dataloader传入的index为int型作为dataset下标;

自定义的sampler则可以传入自定义item,sampler定义参照scepter.modules.utils.sampler;

  • 具体读取单条数据由_get()方法传入index参数实现;
  • 输出经由pipeline转换后的数据(如果有pipeline);

function worker_init_fn()

Parameters(输入参数):

(worker_id, num_workers = 1)

worker_id为分布式训练中工作节点id;

dataloader一次性创建num_workers个工作进程;

  • 用于初始化文件读取系统和设置多卡worker对应参数;

function _get()

Parameters(输入参数):index(由__getitem__()传入)

抽象方法,需要由自定义dataset继承具体实现,用于根据index读取batch中每条数据;


基础用法

子类注册:

from scepter.modules.data.dataset import BaseDataset
from scepter.modules.data.dataset import DATASETS


@DATASETS.register_class()
class XxxxxDataset(BaseDataset):
    def __init__(self, cfg, logger=None):
        super(XxxxxDataset, self).__init__(cfg, logger=logger)

实际启动dataset用法:

if self.cfg.have("TRAIN_DATA"):
    train_data = DATASETS.build(self.cfg.TRAIN_DATA, logger=self.logger)

详见scepter/modules/solver/diffusion_solver.py,具体参数定义在配置yaml中TRAIN_DATA/EVAL_DATA/TEST_DATA模块下;