我要提问
ARTICLE DETAIL

资讯详情

前沿编程新知与开发实战干货的深度解读。

TensorFlow Models NLP Losses 模块解析:weighted_sparse_categorical_crossentropy_loss 的实现与实战

TensorFlow Models NLP Losses 模块解析:weighted_sparse_categorical_crossentropy_loss 的实现与实战 TensorFlow Models NLP Losses 模块解析weighted_sparse_categorical_crossentropy_loss 的实现与实战【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models本指南基于 TensorFlow Models 仓库中official/nlp/modeling/losses模块的官方文档展开完整覆盖其唯一的损失函数weighted_sparse_categorical_crossentropy_loss的语义、参数约束与边界行为并结合该损失函数的源码实现与测试用例讲清逐样本损失如何加权平均成批级标量这一 NLP 训练Masked LM、分类、QA 等任务中的核心机制。读完后你能准确使用该损失函数理解其张量秩校验、Keras 标签维度处理与全零权重时的数值安全行为。1. Losses 模块定位official/nlp/modeling/losses/README.md对模块的定位非常明确Losses contains common loss computation used in NLP tasks.weighted_sparse_categorical_crossentropy_losscomputes per-batch sparse categorical crossentropy loss.也就是说该子包是official/nlp/modeling这一 TensorFlow 2.x Keras 版 NLP 建模框架BERT/Electra 等编码器中专门存放跨任务通用损失计算的位置当前模块内只有一个损失函数但它是 Masked LM、句级分类、问答等任务共用的基础构件。目录结构如下文件作用README.md模块说明文档init.py对外导出别名weighted_sparse_categorical_crossentropy_lossweighted_sparse_categorical_crossentropy.py损失函数核心实现weighted_sparse_categorical_crossentropy_test.py单元测试覆盖 3D 输入、权重掩码、秩校验与数值基准official/nlp/modeling/README.md所在的建模框架通过official/nlp/modeling/__init__.py将losses子包注册为顶层可导入命名空间因此在任务侧可以直接这样引入from official.nlp.modeling.losses import weighted_sparse_categorical_crossentropy # 或使用 __init__.py 中导出的别名 from official.nlp.modeling import losses losses.weighted_sparse_categorical_crossentropy_loss2. loss() 函数签名与参数语义核心实现位于 weighted_sparse_categorical_crossentropy.py函数签名为def loss(labels, predictions, weightsNone, from_logitsFalse):各参数的精确语义摘自源码 docstring并结合实现补充参数类型/形状说明labels整型张量秩 rank(predictions) - 1类别索引集合取值范围0 ~ vocab_size - 1实现内会先tf.cast到int32predictions浮点张量形状如(batch, num_positions, vocab_size)或(batch, num_classes)模型输出。默认应已应用 softmax若传的是 logits必须显式设置from_logitsTrueweights与labels同形状秩相同的张量可选逐样本权重为None时对全部样本取平均from_logitsbool默认Falsepredictions是否为未归一化 logits返回值是一个标量损失无weights时是逐样本损失的均值有weights时是加权平均。当传入张量的秩不满足约束时抛出RuntimeError。3. 三个内部实现细节源码除了主函数外还有两个关键辅助函数理解它们是正确使用该损失的前提。3.1 标签 squeeze兼容 Keras 的多余内维def _adjust_labels(labels, predictions): Adjust the labels tensor by squeezing it if needed. labels tf.cast(labels, tf.int32) if len(predictions.shape) len(labels.shape): labels tf.squeeze(labels, [-1]) return labels, predictions源码注释解释了动机Keras 核心 API 会在标签张量末尾附加一个多余的内维例如本应为(batch, num_positions)的标签变成(batch, num_positions, 1)。该函数检测到labels与predictions秩相同时自动压掉最后一维使函数对原始 NumPy 标签和Keras 包装后的标签都能工作。测试用例test_mismatched_predictions_and_labels_ranks_squeezes专门验证了(batch, 1)的标签传入(batch, 10)的 predictions 时 squeeze 成功。3.2 秩校验def _validate_rank(labels, predictions, weights): if weights is not None and len(weights.shape) ! len(labels.shape): raise RuntimeError( (Weight and label tensors were not of the same rank. ...) if (len(predictions.shape) - 1) ! len(labels.shape): raise RuntimeError( (Weighted sparse categorical crossentropy expects labels to have a rank of one less than predictions. ...))两条硬约束labels的秩必须恰好比predictions少 1稀疏标签是类别索引不带类别维若提供weights其秩必须与labels一致。测试test_mismatched_weights_and_labels_ranks_fail用weights形状(batch,)对labels形状(batch, 10)的组合验证了第 2 条约束断言抛出包含of the same rank的RuntimeError。3.3 加权平均与数值安全example_losses tf_keras.losses.sparse_categorical_crossentropy( labels, predictions, from_logitsfrom_logits) if weights is None: return tf.reduce_mean(example_losses) weights tf.cast(weights, predictions.dtype) return tf.math.divide_no_nan( tf.reduce_sum(example_losses * weights), tf.reduce_sum(weights))三个要点逐样本损失委托给tf_keras.losses.sparse_categorical_crossentropyfrom_logits原样透传加权平均采用sum(loss * w) / sum(w)而非mean(loss * w)等价于只对权重非零的样本求平均符合 Masked LM 中仅对真实被遮盖位置计损失、其余填充位置权重为 0 的典型用法使用tf.math.divide_no_nan当weights全为 0 时分母为 0函数安全地返回 0 而不是 NaN。测试test_loss_weights_3d_input构造了一个全零权重张量并断言assertAllClose(0, weighted_loss_data)印证了这一行为。4. 测试用例解读从 Masked LM 到二分类的完整验证链weighted_sparse_categorical_crossentropy_test.py 中的ClassificationLossTest提供了该损失从模型到标量的端到端用法值得逐条对照。4.1 用 BertEncoder MaskedLM 构造真实 3D 输出测试首先搭建了一个最小语言模型test_loss_3d_input第 25–89 行xformer_stack networks.BertEncoder( vocab_sizevocab_size, # 100 num_layers1, sequence_lengthsequence_length, # 32 hidden_sizehidden_size, # 64 num_attention_heads4, ) word_ids tf_keras.Input(shape(sequence_length,), dtypetf.int32) mask tf_keras.Input(shape(sequence_length,), dtypetf.int32) type_ids tf_keras.Input(shape(sequence_length,), dtypetf.int32) _ xformer_stack([word_ids, mask, type_ids]) test_layer layers.MaskedLM( embedding_tablexformer_stack.get_embedding_table(), outputoutput) lm_input_tensor tf_keras.Input(shape(sequence_length, hidden_size)) masked_lm_positions tf_keras.Input(shape(num_predictions,), dtypetf.int32) output test_layer(lm_input_tensor, masked_positionsmasked_lm_positions)随后对batch_size3的随机输入取模型输出计算损失labels np.random.randint(vocab_size, size(batch_size, num_predictions)) weights np.random.randint(2, size(batch_size, num_predictions)) per_example_loss_data weighted_sparse_categorical_crossentropy.loss( predictionsoutput_data, labelslabels, weightsweights) expected_shape [] # Scalar self.assertEqual(expected_shape, per_example_loss_data.shape.as_list())这确认了两个使用契约predictions是 3 维的(batch, num_predictions, vocab_size)MaskedLM 对每个被预测位置输出一个 vocab 维向量labels与weights是 2 维的(batch, num_predictions)最终输出是形状为[]的标量且对随机数据非零。4.2 数值基准回归测试两个test_legacy_*_compatibility用例用固定输入锁死了数值正确性rtol1e-3# Masked LM 场景vocab_size52 个预测位置仅 1 个位置权重为 1 weights np.array([[1, 0], [0, 0], [0, 0]]) loss_data weighted_sparse_categorical_crossentropy.loss( predictionsoutput_data, labelslabels, weightsweights, from_logitsTrue) expected_loss_data 1.2923441 self.assertAllClose(expected_loss_data, loss_data, rtol1e-3)以及纯分类场景batch_size2, num_classes32D 输出 weightsNoneloss_data weighted_sparse_categorical_crossentropy.loss( predictionsoutput_data, labelslabels, weightsNone, from_logitsTrue) expected_loss_data 6.4222这两组用例说明同一函数同时支撑多位置预测的序列级损失Masked LM/QA与单标签分类损失两种形态且注释表明它们是为重构期间保证与旧版计算一致性而建立的基准。5. 实际任务中的使用方式与注意事项official/nlp下的具体任务如 masked_lm.py在 Keras 化重构后对逐样本交叉熵多采用内联调用tf_keras.losses.sparse_categorical_crossentropy的写法再自行按 mask 加权而本模块的loss()则提供了封装好的校验 squeeze 加权平均一体化入口。基于源码与测试使用时的可复制要点如下import numpy as np import tensorflow as tf from official.nlp.modeling.losses import ( weighted_sparse_categorical_crossentropy as wsce) # 场景Masked LM 预测logits 输出 logits tf.random.normal((3, 21, 100)) # (batch, num_predictions, vocab_size) labels np.random.randint(100, size(3, 21)) # (batch, num_predictions) mask np.random.randint(2, size(3, 21)) # 仅对真实遮盖位置计损 loss_val wsce.loss( labelslabels, predictionslogits, weightsmask, from_logitsTrue, # predictions 是 logits 时必须置 True ) # loss_val 为标量若 mask 全 0返回 0 而非 NaN注意事项均可由源码直接确认默认假设 softmax 输入from_logits默认为False把 logits 直接传入会得到数值上不正确的损失必须显式传Trueweights 秩必须与 labels 相同逐样本权重不是逐 batch 标量形状错误会直接RuntimeErrorKeras 内联调用时标签可能带多余内维_adjust_labels会自动 squeeze因此在tf_keras.Model的compile(loss...)场景下无需手工处理标签形状全零权重返回 0由divide_no_nan保证训练初期某个 batch 内没有有效遮盖位置时不会产生 NaN。6. 模块在 NLP 建模框架中的位置从源码结构看losses与layers、models、networks、ops共同组成 official/nlp/modeling/README.md 描述的 Keras 化建模层networks.BertEncoder等网络输出句向量或 logitslayers.MaskedLM等任务头映射回词表空间而本模块的损失函数负责把逐位置的预测误差按 mask 聚合为可反传的标量被official/nlp/tasks下的训练入口train.py间接依赖。official/nlp/README.md的 API 索引中也列出了该损失函数对应tfm.nlp.losses.weighted_sparse_categorical_crossentropy_loss入口表明它是面向外部使用者的公开 API 之一。7. 小结official/nlp/modeling/losses当前提供weighted_sparse_categorical_crossentropy_loss一个损失函数语义是逐样本稀疏交叉熵 可选加权平均成批级标量实现上的三个关键设计自动 squeeze Keras 多余标签维、rank(predictions) - 1 rank(labels)的严格秩校验、divide_no_nan带来的全零权重安全返回测试文件提供了 Masked LM3D 输出与二分类2D 输出两套固定数值的回归基准以及真实BertEncoder MaskedLM前向链路验证适用前提是 TensorFlow 2 /tf_keras环境且当predictions为 logits 时必须显式from_logitsTrue。【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表