CRF层需自定义或使用tensorflow-addons;TF≤2.14用tfa.layers.CRF,TF≥2.15需手写,核心是前向算法与Viterbi解码,输入为logits,loss为负对数似然。

CRF层在TensorFlow里必须用tf.keras.layers.Layer自定义
TensorFlow官方Keras没有内置CRF层,直接调用tf.keras.layers.CRF会报AttributeError。你得自己封装CRF逻辑,或依赖第三方库(如tensorflow-addons)。但注意:tensorflow-addons的tfa.layers.CRF只支持TensorFlow 2.8–2.14,且不兼容TF 2.15+(因底层tf.linalg.band_part行为变更)。
推荐做法是:用tfa.layers.CRF(TF ≤2.14)或手写CRF层(TF ≥2.15)。手写时核心是实现前向算法和Viterbi解码,不能只靠tf.nn.softmax——那只是逐帧归一化,不是序列级建模。
- CRF层输入必须是logits(未softmax),shape为
(batch_size, seq_len, num_tags) - 转移矩阵需初始化为
tf.Variable,shape为(num_tags, num_tags),不能是常量 - 训练时loss要包含标签路径得分与所有路径总分之差,即
log_likelihood取负
用tensorflow-addons加载CRF层的典型报错及修复
常见错误是ModuleNotFoundError: No module named 'tfa'或InvalidArgumentError: input must be at least 2-dim。前者说明没装对版本;后者多因输入张量少了一维——CRF层要求输入是3D,而LSTM输出若用了return_sequences=False,只剩2D,必须设为True。
安装命令要严格匹配TF版本:
立即学习“Python免费学习笔记(深入)”;
pip install tensorflow-addons==0.21.0 # for TF 2.11–2.14 pip install tensorflow-addons==0.20.0 # for TF 2.10
模型构建示例:
import tensorflow as tf import tensorflow_addons as tfa <p>model = tf.keras.Sequential([ tf.keras.layers.Embedding(vocab_size, 100), tf.keras.layers.Bidirectional(tf.keras.layers.LSTM(128, return_sequences=True)), tf.keras.layers.Dense(num_tags), # logits, not softmax! tfa.layers.CRF(num_tags) # 自动处理转移矩阵和loss ])
- 必须在
Dense后接CRF,且Dense不能带activation='softmax' -
tfa.layers.CRF默认返回(decoded_sequence, log_likelihood),训练时要用log_likelihood作为loss - 预测时用
model.predict()得到的是decoded_sequence,不是概率分布
手动实现CRF层的关键三步:转移矩阵、前向算法、Viterbi
手写CRF不难,但容易在维度和mask处理上出错。核心不是重写整个算法,而是正确调用tf.linalg.band_part(旧版)或改用tf.where + tf.scatter_nd(新版)构造转移约束;前向算法中每步需用tf.reduce_logsumexp而非tf.reduce_sum,否则数值溢出。
- 转移矩阵初始化建议用
tf.random.normal((num_tags, num_tags), stddev=0.1),避免全零导致梯度消失 - 输入mask必须是布尔型
tf.bool,传入前向函数前要转成tf.float32并广播到logits维度 - Viterbi回溯时,用
tf.TensorArray暂存每步最大路径索引,不能用Python list
关键片段示意(简化版前向):
def crf_log_likelihood(logits, tags, trans, mask):
# logits: [B, T, N], tags: [B, T], trans: [N, N], mask: [B, T]
sequence_lengths = tf.reduce_sum(tf.cast(mask, tf.int32), axis=1)
log_likelihood, _ = tfa.text.crf_log_likelihood(
logits, tags, sequence_lengths, trans
)
return log_likelihood注意:tfa.text.crf_log_likelihood是底层C++实现,比纯tf手写快且稳,只要TF版本兼容就优先用它。
训练时loss和metrics容易混淆的两个点
CRF模型的loss必须是-log_likelihood,不是sparse_categorical_crossentropy。否则模型学不会标签间依赖关系,性能掉点明显(NER任务F1常降5–10%)。同样,评估指标也不能用SparseCategoricalAccuracy——它只比对单个token,忽略序列约束。
- 训练loss设为
lambda y_true, y_pred: -crf_log_likelihood(y_pred, y_true, trans, mask) - 自定义metric需在
update_state中调用tfa.text.crf_decode获得预测标签,再与真实标签逐元素比较 - 如果用
model.compile(loss='sparse_categorical_crossentropy'),哪怕后面接了CRF层,loss也完全失效
最稳妥的方式是把CRF loss写进训练循环里,绕过Keras自动loss机制,尤其在TF ≥2.15时更可控。


















