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
相关产品推荐
相关产品推荐

