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

Keras-contrib CRF层learn_mode参数疑问:join模式为何致Loss为NaN?

Keras-Contrib CRF层learn_mode参数问题解答

我来帮你拆解一下关于Keras-Contrib中CRF层learn_mode参数的这几个问题,结合你做NER任务的场景来解释:

1. 'join'与'marginal'两种learn_mode的区别是什么?

简单来说,这两个参数决定了CRF层的损失计算和训练优化逻辑:

  • join模式:属于联合训练,它直接优化的是整个标签序列的联合对数似然。说白了,模型会学习“给定输入文本,整个正确标签序列出现的概率”,完全利用了序列标签之间的依赖关系(比如NER中B-PER后面大概率跟I-PER),理论上更贴合序列标注的任务本质,但对数值稳定性要求极高。
  • marginal模式:属于边际训练,它优化的是每个位置上对应正确标签的边际对数似然之和。也就是模型会逐个位置计算“这个位置是正确标签的概率”,然后把所有位置的损失加起来。它牺牲了一点对序列全局依赖的建模精度,但胜在数值计算更稳定,训练时不容易出异常。

2. 为何在你的NER任务中'join'模式会导致Loss为NaN?

结合你的模型结构和NER场景,大概率是数值不稳定触发的问题,具体可能有这几个原因:

  • 序列长度与联合概率的冲突:你的模型输入是max_len_doc长度的文本,join模式需要计算整个序列的联合概率,当序列较长时,多个概率项相乘(取对数后是相加)很容易出现数值下溢——比如某个位置的概率趋近于0,取对数后就会变成负无穷,反向传播时梯度直接爆炸,最终Loss变成NaN。
  • 正则化操作的放大影响:你的模型里用了SpatialDropout1D、多层recurrent_dropout,这些正则化会让模型的输出分布更“极端”(比如某些位置的输出概率接近0或1),而join模式对这种极端数值特别敏感,很容易触发数值崩溃。
  • 标签分布的影响:NER任务中通常存在标签不平衡的问题(比如O标签占比极高,某些实体标签很少见),join模式在计算联合似然时,稀有标签对应的项会让整个似然值变得极小,进一步加剧数值下溢的问题。

3. 'marginal'模式为何能正常工作?

刚好对应join模式的痛点,marginal模式的设计天生更耐造:

  • 数值稳定性更强:它是逐个位置计算损失,每个位置的边际概率计算相对独立,不会因为整个序列的联合概率极端而直接崩溃。就算某个位置的概率有异常,也只会影响该位置的损失,不会让全局Loss变成NaN。
  • 对正则化更鲁棒:Dropout这类操作带来的输出波动,在marginal模式下被分散到了每个位置的损失中,不会像join模式那样被放大成整个序列的似然异常。
  • 计算逻辑更“温和”:边际似然是把每个位置的损失简单相加,而联合似然是基于序列依赖的概率乘积(取对数后仍是关联的求和),前者的数值范围更容易控制,反向传播时梯度也更平缓稳定。

附你的模型代码

# input and embedding for words
word_in = Input(shape=(max_len_doc,))
emb_word = Embedding(input_dim=n_words + 2, output_dim=50, input_length=max_len_doc, mask_zero=True)(word_in)
# input and embeddings for characters
char_in = Input(shape=(max_len_doc, max_len_word,))
emb_char = TimeDistributed(Embedding(input_dim=n_chars + 2, output_dim=10, input_length=max_len_word, mask_zero=True))(char_in)
# character LSTM to get word encodings by characters
char_enc = TimeDistributed(LSTM(units=50, return_sequences=False, recurrent_dropout=0.5))(emb_char)
# main LSTM
model_crf = concatenate([emb_word, char_enc])
model_crf = SpatialDropout1D(0.3)(model_crf)
model_crf = Bidirectional(LSTM(units=128, return_sequences=True, recurrent_dropout=0.6))(model_crf)
model_crf = Bidirectional(LSTM(units=128, return_sequences=True, recurrent_dropout=0.3))(model_crf)
model_crf = TimeDistributed(Dense(n_tags, activation="relu"))(model_crf)
crf = CRF(n_tags) # crf = CRF(n_tags, learn_mode='marginal')
out = crf(model_crf)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 03:47:00