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

如何在tf.data.Dataset中用tf.repeat及高阶函数实现NER词片对齐?

NER任务中WordPieces与标签对齐的tf.data实现方案

在NER任务中,使用WordPiece分词后需要将原单词级标签与每个subword对齐(即每个单词的标签重复对应其所有subword),当数据来自tf.data.TextLineDataset时,可通过以下纯TensorFlow原生操作实现,适配流水线且高效简洁。

核心对齐逻辑

假设你的TextLineDataset每条数据已解析为单词级subword列表(RaggedTensor)和单词级标签列表(Tensor),用以下函数完成对齐:

import tensorflow as tf

def align_tags(tokens_per_word, tags):
    # 获取每个单词对应的subword数量(RaggedTensor原生方法,高效无遍历)
    subword_counts = tokens_per_word.row_lengths()
    # 重复标签,次数对应单词的subword数量
    aligned_tags = tf.repeat(tags, subword_counts)
    # 直接提取RaggedTensor的一维token序列
    flat_tokens = tokens_per_word.flat_values
    return flat_tokens, aligned_tags

完整流水线示例

1. 解析文本行数据

如果你的文本文件每行格式为[单词1_subword1,单词1_subword2;单词2_subword1]\t[标签1,标签2](例如"2,5,9;1,3\t5,7"),先编写解析函数将每行转为对应张量:

def parse_line(line):
    # 拆分token段与标签段
    tokens_part, tags_part = tf.strings.split(line, "\t").to_list()
    
    # 解析单词级subword列表:按分号拆分单词,逗号拆分subword并转int
    tokens_per_word = tf.strings.split(tokens_part, ";")
    tokens_per_word = tf.ragged.map_flat_values(
        lambda x: tf.strings.to_number(tf.strings.split(x, ","), tf.int32),
        tokens_per_word
    )
    
    # 解析标签列表:按逗号拆分并转int
    tags = tf.strings.to_number(tf.strings.split(tags_part, ","), tf.int32)
    return tokens_per_word, tags

# 从文本文件加载数据集
dataset = tf.data.TextLineDataset("ner_data.txt").map(parse_line)

2. 应用对齐逻辑并验证

processed_dataset = dataset.map(align_tags)

# 查看前2条处理结果
for tokens, tags in processed_dataset.take(2):
    print("展平后的tokens:", tokens.numpy())
    print("对齐后的tags:", tags.numpy())
    print("---")

关键优势

  • 全程使用TensorFlow原生操作,无Python列表遍历,可序列化适配异步加载、分布式训练;
  • 利用RaggedTensor.row_lengths()高效获取subword数量,性能优于手动计算;
  • flat_values直接提取一维token序列,无需额外reshape操作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 05:07:21