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

PyTorch CTC Loss测试数据表现异常,请求问题排查

CTC Loss计算结果异常问题排查

我在理解CTC Loss工作原理时,生成了几组输入与真值数据,传入CTC Loss函数后得到了不符合预期的结果。

"|" 是vocab[4],代表空白字符

测试示例

测试#1 "精准猜测"

  • 目标词: "acb"、"acc"
  • 预测结果: "a|ccc|b"、"aaa|ccc|cccc"
  • 损失值: tensor([16.9247, 18.9556])

测试#2 "完全匹配,无空白字符"

  • 目标词: "b"、"b"
  • 预测结果: "b"、"b"
  • 损失值: tensor([22.0504, 22.0504])

测试#3 "完全不匹配"

  • 目标词: "bba"、"bbb"
  • 预测结果: "a|ccc|b"、"aaa|ccc|cccc"
  • 损失值: tensor([17.4538, 18.9426])

三次测试的损失值几乎相同,我原本预期前两次测试的损失值较低,第三次测试损失值较高。请问我哪里操作出错了?

代码实现

vocab_test = ['a', 'b', 'c',' ', '|']
vocab_dict_test = {'a':1,'b':2,'c':3,' ':4, '|':5}

vocab_length = 5
batch_size = 2
label_len = 7
input_len = 15

def loss_check(word1, word2, guess1, guess2):

    # convert words into digital torches of size needed
    input_word1 = word1.ljust(label_len)
    guessed_word1 = guess1.ljust(input_len)

    input_word_numbers1 = get_numbers(vocab_dict_test, input_word1)
    guessed_word_numbers1 = get_numbers(vocab_dict_test, guessed_word1)

    input_word2 = word2.ljust(label_len)
    guessed_word2 = guess2.ljust(input_len)

    input_word_numbers2 = get_numbers(vocab_dict_test, input_word2)
    guessed_word_numbers2 = get_numbers(vocab_dict_test, guessed_word2)

    print('words:', '"' + word1 + '" / "' + word2 + '"', 'converted:', input_word_numbers1, input_word_numbers2)
    print('guesses:', '"' + guess1 + '" / "' + guess2 + '"', 'converted:', guessed_word_numbers1, guessed_word_numbers2)

    # special torches for loss func
    input_len_size = torch.IntTensor([input_len] * batch_size).to(device)
    label_len_size = torch.IntTensor([label_len] * batch_size).to(device)

    truth = torch.from_numpy(np.array([input_word_numbers1,input_word_numbers2])).float()
    #print(truth)
    #print(truth.size())

    logits = [torch.from_numpy(np.array(guessed_word_numbers1)),torch.from_numpy(np.array(guessed_word_numbers2))]
    logits_converted = [[[1 if (logits[j][h] == i+1) else 0 for i in range(vocab_length)] for h in range(input_len)] for j in range(batch_size)]
    logits_converted = torch.from_numpy(np.array(logits_converted)).float()
    #print(logits_converted)
    Softmax = nn.LogSoftmax(dim=2)
    logits_log = Softmax(logits_converted)
    #print(logits_log)

    logits_log_formatted = logits_log.transpose(1,0)
    #print(logits_log_formatted)
    #print(logits_log_formatted.size())
    loss_fn = nn.CTCLoss(blank=4, zero_infinity=True, reduction='none').to(device)

    #log_probs, targets, input_lengths, target_lengths, self.blank,
    loss = loss_fn(logits_log_formatted, truth, input_len_size, label_len_size)

    return loss


print(loss_check('acb', 'acc','a|ccc|b', 'aaa|ccc|cccc'))
print(loss_check('bba', 'bbb','a|ccc|b', 'aaa|ccc|cccc'))
print(loss_check('b', 'b','b', 'b'))

内容的提问来源于stack exchange,提问作者Olga

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 03:17:05