TensorFlow调用Model.fit()训练Tacotron2时张量形状类型报错
TensorFlow实现Tacotron2调用Model.fit()时动态维度报错问题
问题复现
- 基于TensorFlow自实现Tacotron2模型,手动传入批次前向计算正常,调用
Model.fit()启动训练时触发类型错误 - 数据集采用
tf.data.Dataset构建,返回格式为((phonemes, mel_spec), (mel_spec, gates)):phonemes:长度可变的音素字符串序列mel_spec:时间维度长度可变、通道数固定为80的二维梅尔声谱图,因教师强制训练机制,同时作为输入和输出gates:长度与mel_spec时间维度一致的一维停止位预测张量
- 因序列长度可变,采用
padded_batch做填充批处理,仅将gates的填充值设为1:
dataset = dataset.padded_batch(batch_size, padding_values=((None, None), (None, 1.)) )
- 手动校验数据集和前向逻辑均正常:
- 打印批次数据,张量形状、填充逻辑符合预期
- 手动拉取批次传入模型执行前向,返回张量形状完全正确,校验代码如下:
x, y = next(iter(dataset.padded_batch(batch_size, padding_values=((None, None), (None, 1.)) ))) mels, gates = tac(x)
- 训练启动代码如下,运行即报错:
dataset = dataset.padded_batch(batch_size, padding_values=((None, None), (None, 1.)) ) optimizer = conf["train"]["optimizer"] epochs = conf["train"]["epochs"] tac.compile(optimizer=optimizer, loss=tac.criterion) tac.fit(dataset, epochs=epochs)
- 报错栈指向模型
call方法中的维度裁剪逻辑:
Tacotron2.py:177 call * crop = mels.shape[2] - mels.shape[2]%self.config["n_frames_per_step"]#max_len must be a multiple of n_frames_per_step TypeError: unsupported operand type(s) for %: 'NoneType' and 'int'
对应的call方法原始实现:
def call(self, batch, training=False): phon, mels = batch x = self.tokenizer(phon) x = self.char_embedding(x) y = self.encoder(x) print(y.shape) crop = mels.shape[2] - mels.shape[2]%self.config["n_frames_per_step"]#max_len must be a multiple of n_frames_per_step mels, gates = self.decoder(y, mels[:,:,:crop]) residual = self.decoder.postnet(mels) mels_post = mels + residual return (mels, mels_post), gates
根因分析
Model.fit()默认会在静态图模式下执行计算图追踪,此时对于长度可变的时间维度,张量的静态形状属性tensor.shape对应维度会返回None(代表该维度动态可变,编译阶段无确定值),直接用这个None和整数做取模运算就会触发类型错误。
手动拉取批次执行前向是在动态图(eager)模式下运行,此时张量所有维度都有确定的实际数值,所以不会触发该问题。
修复方案
涉及动态可变维度的运算,不要使用静态形状属性tensor.shape取值,改用TensorFlow提供的tf.shape(tensor)接口获取运行时的实际维度值,修改call方法中的裁剪逻辑即可:
def call(self, batch, training=False): phon, mels = batch x = self.tokenizer(phon) x = self.char_embedding(x) y = self.encoder(x) # 用tf.shape获取运行时动态维度值,替换静态shape读取 mel_time_len = tf.shape(mels)[2] crop = mel_time_len - mel_time_len % self.config["n_frames_per_step"] mels, gates = self.decoder(y, mels[:, :, :crop]) residual = self.decoder.postnet(mels) mels_post = mels + residual return (mels, mels_post), gates
注意:后续所有涉及动态可变维度的计算、切片操作,都需要使用tf.shape()取运行时值;只有配置固定、不会随输入变化的维度(比如梅尔通道数80、卷积输出通道数、嵌入维度等),才可以直接用tensor.shape读取静态值。
内容的提问来源于stack exchange,提问作者purple
相关产品推荐
相关产品推荐

