在MindSpore2.3版本中,使用LSTM模型做藏头诗的生成工作,模型训练过程出现BUG。
收藏回复举报
在MindSpore2.3版本中,使用LSTM模型做藏头诗的生成工作,模型训练过程出现BUG。
t('forum.solved') 已解决
新人帖
发表于2024-11-06 17:49:06
0 查看

LSTM模型结构如下:

"""LSTM.""" 

from mindspore import nn 

from mindspore.ops import operations as P 

 

class SentimentNet(nn.Cell): 

    """Sentiment network structure.""" 

 

    def __init__(self, 

                 vocab_size, 

                 embed_size, 

                 num_hiddens, 

                 num_layers, 

                 bidirectional, 

                 num_classes, 

                 weight, 

                 batch_size): 

        super(SentimentNet, self).__init__() 

        self.bidirectional = bidirectional 

        # Mapp words to vectors 

        self.embedding = nn.Embedding(vocab_size, 

                                      embed_size, 

                                      False) 

                                      # embedding_table=weight) 

        self.embedding.embedding_table.requires_grad = False 

        self.trans = P.Transpose() 

        self.perm = (1, 0, 2) 

 

        self.encoder = nn.LSTM(input_size=embed_size, 

                               hidden_size=num_hiddens, 

                               num_layers=num_layers, 

                               has_bias=True, 

                               bidirectional=bidirectional, 

                               dropout=0.0) 

 

        self.concat = P.Concat(1) 

        self.squeeze = P.Squeeze(axis=0) 

        if self.bidirectional: 

            self.decoder = nn.Dense(num_hiddens * 4, num_classes) 

        else: 

            self.decoder = nn.Dense(num_hiddens * 2, num_classes) 

 

    def construct(self, inputs): 

        # input:(64,500,300) 

        embeddings = self.embedding(inputs) 

        embeddings = self.trans(embeddings, self.perm) 

        output, _ = self.encoder(embeddings) 

        # states[i] size(64,200)  -> encoding.size(64,400) 

        encoding = self.concat((self.squeeze(output[0:1:1]), self.squeeze(output[499:500:1]))) 

        outputs = self.decoder(encoding) 

        print(outputs) 

        return outputs 

mindrecord的数据格式如下:

import mindspore.dataset as ds 

import mindspore.nn as nn 

import mindspore.ops as ops 

for data in ds_train.create_dict_iterator(output_numpy=True): 

    features = data['feature'] 

    labels = data['label'] 

    print(f"Features shape: {features.shape}") 

    print(f"Labels shape: {labels.shape}") 

    break  # 只打印第一批数据

Features shape: (8, 500)
Labels shape: (8, 500)

LSTM模型loss如下:

 

import mindspore as ms 

from mindspore.nn.loss.loss import _Loss 

#定义loss 

class NLLLoss(_Loss): 

    ''' 

       NLLLoss function 

    ''' 

    def __init__(self, reduction='mean'): 

        super(NLLLoss, self).__init__(reduction) 

        self.reduce_sum = P.ReduceSum() 

        self.one_hot = P.OneHot()#标签是稀疏的,所以用one-hot转成向量再进行计算 

    def construct(self, prob, label): 

        label_one_hot = self.one_hot(label, F.shape(prob)[-1], F.scalar_to_array(1.0), ops.scalar_to_array(0.0)) 

        loss = self.reduce_sum(-1.0 * prob * label_one_hot, (1,)) 

        return self.get_loss(loss) 

class LSTMWithLossCell(nn.Cell):  

 

    def __init__(self, network): 

        super(LSTMWithLossCell, self).__init__() 

        self.network = network 

        self.loss = NLLLoss() 

        self.squeeze = P.Squeeze() 

        self.add = P.AddN() 

    def construct(self, x, y):         

        logits,_ = self.network(x) 

        loss_total = () 

        self.text_len = len(y[0]) 

        for i in range(self.text_len): 

            loss = self.loss(self.squeeze(logits[i, ::, ::]), y[:,i]) 

            loss_total += (loss,) 

        loss = self.add(loss_total) / self.text_len 

        return loss 

opt = nn.Momentum(network.trainable_params(), lr, cfg.momentum) 

loss_cb = LossMonitor() 

network.set_jit_config(JitConfig(jit_level="O2")) 

network = WithLossCell(network, cfg) 

optimizer = nn.Adam(network.trainable_params(), learning_rate=cfg.learning_rate, beta1=0.9, beta2=0.98) 

model = Model(network, optimizer=optimizer) 

print("============== Starting Training ==============") 

config_ck = CheckpointConfig(save_checkpoint_steps=cfg.save_checkpoint_steps, 

                             keep_checkpoint_max=cfg.keep_checkpoint_max) 

ckpoint_cb = ModelCheckpoint(prefix="lstm", directory=cfg.ckpt_path, config=config_ck) 

time_cb = TimeMonitor(data_size=ds_train.get_dataset_size()) 

cb = [time_cb, loss_cb, ckpoint_cb]

rank = 0 

device_num = 1 

ds_train = lstm_create_dataset(cfg.preprocess_path, cfg.batch_size, device_num=device_num, rank=rank) 

model.train(cfg.num_epochs, ds_train, callbacks=cb, dataset_sink_mode=False) 

print("============== Training Success ==============")

最终报错如下:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In[123], line 4
      2 device_num = 1
      3 ds_train = lstm_create_dataset(cfg.preprocess_path, cfg.batch_size, device_num=device_num, rank=rank)
----> 4 model.train(cfg.num_epochs, ds_train, callbacks=cb, dataset_sink_mode=False)
      5 print("============== Training Success ==============")

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/train/model.py:1082, in Model.train(self, epoch, train_dataset, callbacks, dataset_sink_mode, sink_size, initial_epoch)
   1079 if callbacks:
   1080     self._check_methods_for_custom_callbacks(callbacks, "train")
-> 1082 self._train(epoch,
   1083             train_dataset,
   1084             callbacks=callbacks,
   1085             dataset_sink_mode=dataset_sink_mode,
   1086             sink_size=sink_size,
   1087             initial_epoch=initial_epoch)
   1089 # When it's distributed training and using MindRT,
   1090 # the node id should be reset to start from 0.
   1091 # This is to avoid the timeout when finding the actor route tables in 'train' and 'eval' case(or 'fit').
   1092 if _enable_distributed_mindrt():

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/train/model.py:115, in _save_final_ckpt.<locals>.wrapper(self, *args, **kwargs)
    113         raise e
    114 else:
--> 115     func(self, *args, **kwargs)

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/train/model.py:630, in Model._train(self, epoch, train_dataset, callbacks, dataset_sink_mode, sink_size, initial_epoch, valid_dataset, valid_frequency, valid_dataset_sink_mode)
    628 self._check_reuse_dataset(train_dataset)
    629 if not dataset_sink_mode:
--> 630     self._train_process(epoch, train_dataset, list_callback, cb_params, initial_epoch, valid_infos)
    631 elif context.get_context("device_target") == "CPU":
    632     logger.info("The CPU cannot support dataset sink mode currently."
    633                 "So the training process will be performed with dataset not sink.")

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/train/model.py:932, in Model._train_process(self, epoch, train_dataset, list_callback, cb_params, initial_epoch, valid_infos)
    930 list_callback.on_train_step_begin(run_context)
    931 self._check_network_mode(self._train_network, True)
--> 932 outputs = self._train_network(*next_element)
    933 cb_params.net_outputs = outputs
    934 if self._loss_scale_manager and self._loss_scale_manager.get_drop_overflow_update():

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:693, in Cell.__call__(self, *args, **kwargs)
    691 except Exception as err:
    692     _pynative_executor.clear_res()
--> 693     raise err
    695 return output

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:689, in Cell.__call__(self, *args, **kwargs)
    687 try:
    688     _pynative_executor.new_graph(self, *args, **kwargs)
--> 689     output = self._run_construct(args, kwargs)
    690     _pynative_executor.end_graph(self, output, *args, **kwargs)
    691 except Exception as err:

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:477, in Cell._run_construct(self, cast_inputs, kwargs)
    475     output = self._shard_fn(*cast_inputs, **kwargs)
    476 else:
--> 477     output = self.construct(*cast_inputs, **kwargs)
    478 if self._enable_forward_hook:
    479     output = self._run_forward_hook(cast_inputs, output)

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/wrap/cell_wrapper.py:418, in TrainOneStepCell.construct(self, *inputs)
    416 def construct(self, *inputs):
    417     if not self.sense_flag:
--> 418         return self._no_sens_impl(*inputs)
    419     loss = self.network(*inputs)
    420     sens = F.fill(loss.dtype, loss.shape, self.sens)

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/wrap/cell_wrapper.py:433, in TrainOneStepCell._no_sens_impl(self, *inputs)
    431 def _no_sens_impl(self, *inputs):
    432     """construct implementation when the 'sens' parameter is passed in."""
--> 433     loss = self.network(*inputs)
    434     grads = self.grad_no_sens(self.network, self.weights)(*inputs)
    435     grads = self.grad_reducer(grads)

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:693, in Cell.__call__(self, *args, **kwargs)
    691 except Exception as err:
    692     _pynative_executor.clear_res()
--> 693     raise err
    695 return output

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:689, in Cell.__call__(self, *args, **kwargs)
    687 try:
    688     _pynative_executor.new_graph(self, *args, **kwargs)
--> 689     output = self._run_construct(args, kwargs)
    690     _pynative_executor.end_graph(self, output, *args, **kwargs)
    691 except Exception as err:

File /usr/local/python3.9.2/lib/python3.9/site-packages/mindspore/nn/cell.py:477, in Cell._run_construct(self, cast_inputs, kwargs)
    475     output = self._shard_fn(*cast_inputs, **kwargs)
    476 else:
--> 477     output = self.construct(*cast_inputs, **kwargs)
    478 if self._enable_forward_hook:
    479     output = self._run_forward_hook(cast_inputs, output)

TypeError: construct() missing 1 required positional argument: 'label'

我要发帖子