You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 08:09:15