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

9.9 KiB
Raw Blame History

快速开始 (Quick Start)

本章节以stable diffusion v1.5为例,介绍如何从零开始构建一个网络结构,以及基于该结构训练一个模型和测试该模型。

1. 使用network类定义模型结构

network类包含了定义、训练和测试模型的方法,我们首先初始化一个network类,然后在里面定义所需的autoencoder, unet, embedder, 以及loss子模块。

@MODELS.register_class()
class LatentDiffusion(TrainModule):
    def __init__(self, cfg, logger=None):
        super().__init__(cfg, logger=logger)
        self.model_config = cfg.DIFFUSION_MODEL
        self.first_stage_config = cfg.FIRST_STAGE_MODEL
        self.cond_stage_config = cfg.COND_STAGE_MODEL
        self.loss_config = cfg.get('LOSS', None)

        self.model = BACKBONES.build(self.model_config, logger=self.logger)
        self.first_stage_model = MODELS.build(self.first_stage_config,
                                              logger=self.logger)
        self.cond_stage_model = EMBEDDERS.build(self.cond_stage_config,
                                                logger=self.logger)
        if self.loss_config:
            self.loss = LOSSES.build(self.loss_config, logger=self.logger)

        # 其他变量和模块定义

2. 实现自定义network类的训练和测试方法

每个network类依赖forward_train和forward_test方法定义自己的训练和测试流程,在sd1.5中,forward_train对采样时刻t进行噪声预测以及进行loss计算

 def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
    x_start = self.encode_first_stage(image, **kwargs)
    t = torch.randint(0,
                        self.num_timesteps, (x_start.shape[0], ),
                        device=x_start.device).long()
    context = {}
    if prompt and self.cond_stage_model:
        zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist()
        prompt = [
            self.train_n_prompt if zeros[idx] else p
            for idx, p in enumerate(prompt)
        ]

        with torch.autocast(device_type='cuda', enabled=False):
            context = self.encode_condition(
                self.tokenizer(prompt).to(we.device_id))

    loss = self.diffusion.loss(x0=x_start,
                                t=t,
                                model=self.model,
                                model_kwargs={'cond': context},
                                noise=noise)
    loss = loss.mean()
    ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
    return ret

forward_test函数用于推理阶段执行完整的图像去噪过程

@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
                    prompt=None,
                    n_prompt=None,
                    sampler='ddim',
                    sample_steps=50,
                    seed=2023,
                    guide_scale=7.5,
                    guide_rescale=0.5,
                    discretization='trailing',
                    run_train_n=True,
                    **kwargs):
    g = torch.Generator(device=we.device_id)
    seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
    g.manual_seed(seed)
    num_samples = len(prompt)

    n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
    assert isinstance(prompt, list) and \
            isinstance(n_prompt, list) and \
            len(prompt) == len(n_prompt)

    context = self.encode_condition(self.tokenizer(prompt).to(
        we.device_id),  method='encode_text')
    null_context = self.encode_condition(self.tokenizer(n_prompt).to(
        we.device_id), method='encode_text')

    width, height = 512, 512
    noise = self.noise_sample(num_samples, width // self.size_factor,
                                height // self.size_factor, g)
    # UNet use input n_prompt
    samples = self.diffusion.sample(solver=sampler,
                                    noise=noise,
                                    model=self.model,
                                    model_kwargs=[{
                                        'cond': context
                                    }, {
                                        'cond': null_context
                                    }],
                                    steps=sample_steps,
                                    guide_scale=guide_scale,
                                    guide_rescale=guide_rescale,
                                    discretization=discretization,
                                    show_progress=True,
                                    seed=seed,
                                    condition_fn=None,
                                    clamp=None,
                                    percentile=None,
                                    t_max=None,
                                    t_min=None,
                                    discard_penultimate_step=None,
                                    return_intermediate=None,
                                    **kwargs)
    x_samples = self.decode_first_stage(samples).float()
    x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)

    outputs = list()
    for p, np, img in zip(prompt, n_prompt, x_samples):
        one_tup = {'prompt': p, 'n_prompt': np, 'image': img}
        outputs.append(one_tup)

    return outputs

3. 子模块注册

在实现完network类之后,需要确保network类中用到的所有子模块都已完成注册。以sd1.5中的embedder为例,为了能在network的初始化方法中实例化该embedder,我们需要先实现该embedder类,并注册到scepter中

@EMBEDDERS.register_class()
class FrozenCLIPEmbedder(BaseEmbedder):
    def __init__(self, cfg, logger=None):
        super().__init__(cfg, logger=logger)

        with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL,
                                     wait_finish=True) as local_path:
            self.tokenizer = CLIPTokenizer.from_pretrained(local_path)
            self.transformer = CLIPTextModel.from_pretrained(local_path)

        self.use_grad = cfg.get('USE_GRAD', False)
        self.freeze_flag = cfg.get('FREEZE', True)
        if self.freeze_flag:
            self.freeze()

        self.max_length = cfg.get('MAX_LENGTH', 77)
        self.layer = cfg.get('LAYER', 'last')
        self.layer_idx = cfg.get('LAYER_IDX', None)
        self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False)
        assert self.layer in self.LAYERS
        if self.layer == 'hidden':
            assert self.layer_idx is not None
            assert 0 <= abs(self.layer_idx) <= 12


    def encode_text(self,
                    tokens,
                    tokenizer=None,
                    append_sentence_embedding=False):
        # 定义一些需要的方法
        pass

    @staticmethod
    def get_config_template():
        return dict_to_yaml('MODELS',
                            __class__.__name__,
                            FrozenCLIPEmbedder.para_dict,
                            set_name=True)

4. Solver注册

solver类中封装了训练和测试一个network所需要的完整流程。以sd1.5为例,注册一个训练sd1.5模型的solver需要:

  1. 创建一个或多个数据加载器(data_loader), 对应solver中的construct_data方法;
  2. 实例化一个模型,对应solver中的construct_model方法
  3. (可选)定义度量指标,对应solver中的construct_metrics方法
  4. (可选)定义训练和测试用到的一些钩子(HOOKS),比如模型保存,预训练参数加载,日志打印,等等。
@SOLVERS.register_class()
class LatentDiffusionSolver(BaseSolver):
    def set_up(self):
        self.construct_data()
        self.construct_model()
        self.construct_metrics()
        self.model_to_device()
        self.init_opti()

    def load_checkpoint(self, checkpoint):
        # 这里定义加载模型的指令

    def save_checkpoint(self):
        # 这里定义保存模型的指令

    def solve(self):
        # 入口函数,根据数据类型选择执行训练或测试
        self.before_solve()
        if 'train' in self._mode_set:
            self.run_train()
        if 'test' in self._mode_set:
            self.run_test()
        self.after_solve()

    def run_train(self):
        # 模型训练

    def run_eval(self):
        # 模型验证

    def run_test(self):
        # 模型测试

5. 训练/测试超参数定义

在注册完所需的各个组件(包括但不限于BACKBONE, NETWORK, EMBEDDER, SOLVER, METRIC)后,需要对其中用到的一些超参数进行设置,scepter使用yaml文件定义各模块超参数,具体参考scepter/modules/examples/sd15/sd15_512_full.yaml

6. 模型训练

通过指定--cfg参数加载所需的yaml文件,完成训练或批量测试的操作

多机多卡训练
# 基于spawn方式,为默认模式
export CUDA_VISIBLE_DEVICES=0,1,2,3
export WORLD_SIZE=1
python  -W ignore  scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml
# 基于原生的pytorch引擎 torchrun模式
torchrun --nproc_per_node 4  scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun
# 基于pytorch_lightning引擎, ENV.USE_PL需要设置为true
python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml
# 单卡训练
# 基于spawn方式,为默认模式
export CUDA_VISIBLE_DEVICES=0
export WORLD_SIZE=1
python -W ignore  scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml
# 基于原生的pytorch引擎
python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun
# 基于pytorch_lightning引擎, ENV.USE_PL需要设置为true
python scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml

7. 模型推理

单次推理可以通过自定义run_inference.py文件来实现,参照sd1.5的推理方法

python -W ignore scepter/run_inference.py --prompt "a woman" --n_prompt "" --num_samples 4 --pretrained_model "path/to/your/pretrained/model"