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

如何用TensorFlow按指定整数拆分序列?字符级模型嵌入拓展求助

嘿,这个需求我在做字符级NLP模型的时候刚好碰到过!要在TensorFlow里按指定整数(这里是代表空格的2)拆分序列,还要把分隔符本身也作为单独的子序列保留,其实用TF的基础算子就能搞定,给你一步步拆解实现思路,直接能用在你的模型里~

单样本的实现代码

先从单个张量的情况入手,假设你的输入是一个1D张量:

import tensorflow as tf

# 你的输入张量
input_tensor = tf.constant([3,5,2,3,2,1,1,1], dtype=tf.int32)
separator = 2  # 代表空格的整数

# 第一步:找到所有分隔符的索引位置
sep_indices = tf.squeeze(tf.where(input_tensor == separator), axis=1)

# 第二步:构造拆分的边界长度
split_boundaries = []
prev_pos = 0
for idx in sep_indices:
    # 添加分隔符前的片段长度
    split_boundaries.append(idx - prev_pos)
    # 添加分隔符本身的长度(固定为1)
    split_boundaries.append(1)
    prev_pos = idx + 1
# 添加最后一段的长度
split_boundaries.append(tf.shape(input_tensor)[0] - prev_pos)

# 过滤掉可能出现的0长度(比如序列开头/结尾就是分隔符的情况)
split_boundaries = tf.boolean_mask(split_boundaries, split_boundaries > 0)

# 第三步:拆分张量
split_tensors = tf.split(input_tensor, split_boundaries)

运行后split_tensors就是你要的结果:[<tf.Tensor: shape=(2,), dtype=int32, numpy=array([3, 5])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([2])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([3])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([2])>, <tf.Tensor: shape=(3,), dtype=int32, numpy=array([1, 1, 1])>]

批量数据的处理(适配模型训练)

如果是处理批量输入(比如模型的batch数据),因为每个样本拆分后的子序列长度不一致,推荐用RaggedTensor来存储结果,后续和嵌入层交互会更顺畅:

def split_with_separator(input_seq, separator):
    sep_indices = tf.squeeze(tf.where(input_seq == separator), axis=1)
    split_boundaries = []
    prev_pos = 0
    for idx in sep_indices:
        split_boundaries.append(idx - prev_pos)
        split_boundaries.append(1)
        prev_pos = idx + 1
    split_boundaries.append(tf.shape(input_seq)[0] - prev_pos)
    split_boundaries = tf.boolean_mask(split_boundaries, split_boundaries > 0)
    return tf.split(input_seq, split_boundaries)

# 批量输入示例
batch_input = tf.constant([
    [3,5,2,3,2,1,1,1],
    [2,4,2,5],
    [6,7,8]  # 没有分隔符的情况
], dtype=tf.int32)

# 对批量中的每个样本应用拆分函数,输出RaggedTensor
batch_split_result = tf.map_fn(
    lambda x: split_with_separator(x, 2),
    batch_input,
    fn_output_signature=tf.RaggedTensorSpec(shape=[None], dtype=tf.int32)
)

适配你的字符级模型场景

拆分后,每个子序列对应一个“单词”(包括空格分隔符),你可以:

  • 对每个子序列计算单词嵌入(不管是用预训练的还是自己训练的嵌入层)
  • 把单词嵌入广播到子序列的每个字符位置
  • 和字符本身的嵌入拼接,得到带单词特征的增强型字符嵌入

这样就完美满足了你“拼接字符所属单词的嵌入”的需求~

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 08:23:30