如果原代码是pytorch,训练的时候是通过循环进行的,循环里操作比较复杂,在mindspore中也想通过循环进行训练,而不是model.train(..., ... ,)这个mindspore函数,这可以吗?
收藏回复举报
如果原代码是pytorch,训练的时候是通过循环进行的,循环里操作比较复杂,在mindspore中也想通过循环进行训练,而不是model.train(..., ... ,)这个mindspore函数,这可以吗?
t('forum.solved') 已解决
新人帖
发表于2023-04-21 18:00:49
0 查看

如果原代码是pytorch,训练的时候是通过循环进行的 (例如 for epoch in range(self.load_epoch + 1, cfg.SOLVER.MAX_EPOCH)),循环里操作比较复杂,在mindspore中也想通过循环进行训练,而不是model.train(..., ... ,)这个mindspore函数,这可以吗?如果可以希望给出例子或者案例链接。

下面是pytorch的训练代码:

# 模型训练过程
    def train(self):
        self.model.train()
        self.optim.zero_grad()

        iteration = self.load_iteration + 1
        # Epoch迭代
        for epoch in range(self.load_epoch + 1, cfg.SOLVER.MAX_EPOCH):
            print(str(self.optim.get_lr()))
            if epoch >= cfg.TRAIN.REINFORCEMENT.START:
                self.rl_stage = True
            # 设置DataLoader
            self.setup_loader(epoch)
            running_loss = .0
            running_reward_baseline = .0
            # 每一个Epoch内部Iteration迭代
            with tqdm.tqdm(desc='Epoch %d - train' % epoch, unit='it', total=len(self.training_loader)) as pbar:
                for _, (indices, input_seq, target_seq, gv_feat, att_feats, att_mask) in enumerate(
                        self.training_loader):
                    input_seq = input_seq.cuda()
                    target_seq = target_seq.cuda()
                    gv_feat = gv_feat.cuda()
                    att_feats = att_feats.cuda()
                    att_mask = att_mask.cuda()

                    kwargs = self.make_kwargs(indices, input_seq, target_seq, gv_feat, att_feats, att_mask)
                    # 1、计算模型损失(XE训练 或 SCST训练)
                    loss, loss_info = self.forward(kwargs)
                    # 2、梯度清零(清空过往梯度)
                    self.optim.zero_grad()
                    # 3、计算新梯度及梯度裁剪
                    loss.backward()  # 非混合精度训练
                    utils.clip_gradient(self.optim.optimizer, self.model,
                                        cfg.SOLVER.GRAD_CLIP_TYPE, cfg.SOLVER.GRAD_CLIP)
                    # 4、权重更新
                    self.optim.step()  # 非混合精度训练
                    # 5、(XE)、优化器lr更新(用于XE训练),在SCST时不起作用
                    self.optim.scheduler_step('Iter')  # TODO ms中会自动更新,那该怎么做实现Iter的Step?
                    losses.update(loss.item())
                    self.display(iteration, data_time, batch_time, losses, loss_info)
                    # tqdm 迭代信息更新
                    running_loss += loss.item()
                    if not self.rl_stage:
                        pbar.set_postfix(
                            loss='%.2f' % (running_loss / (_ + 1))
                        )
                    else:
                        running_reward_baseline += loss_info['reward_baseline']
                        pbar.set_postfix(
                            {'loss/r_b': '%.2f/%.2f' % (running_loss / (_ + 1), running_reward_baseline / (_ + 1))}
                        )
                    pbar.update()
                    # print(str(self.optim.get_lr()))
                    iteration += 1

            # 每一个Epoch结束保存模型
            self.save_model(epoch)
            # 模型验证测试,返回的val仅用于SCST训练过程
            val = self.eval(epoch)
            # 4(SCST)、优化器lr更新(用于SCST训练),在XE训练时不起作用
            # 4 (XE)、优化器lr更新,当使用Step学习率策略时作用
            self.optim.scheduler_step('Epoch', val)  # TODO 没看懂,ms是自动更新怎么办?
            self.scheduled_sampling(epoch)

我要发帖子