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

如何在TensorFlow中实现类似itertools.permutations的排列生成功能

TensorFlow实现等价itertools.permutations的排列生成函数

我们需要实现自定义函数tf_permutations,输入一维张量即可输出所有元素的全排列张量,支持如下调用方式:

out = tf_permutations(tf.range(3))

完整实现代码

import tensorflow as tf

def tf_permutations(inputs):
    inputs = tf.convert_to_tensor(inputs)
    # 校验输入为一维张量
    tf.debugging.assert_rank(inputs, 1, message="输入必须为一维张量")
    elem_count = tf.shape(inputs)[0]

    # 边界情况:只有1个元素时直接返回
    if elem_count == 1:
        return tf.expand_dims(inputs, axis=0)
    
    # 递归生成前n-1个元素的全排列
    sub_perms = tf_permutations(inputs[:-1])
    sub_perm_count = tf.shape(sub_perms)[0]
    last_elem = inputs[-1]

    # 将最后一个元素插入到每个子排列的所有可插入位置
    sub_perms_repeat = tf.repeat(sub_perms, repeats=elem_count, axis=0)
    insert_pos = tf.tile(tf.range(elem_count), multiples=[sub_perm_count])

    return tf.map_fn(
        lambda x: tf.experimental.numpy.insert(x[0], x[1], last_elem),
        (sub_perms_repeat, insert_pos),
        fn_output_signature=inputs.dtype
    )

调用测试

执行测试代码:

out = tf_permutations(tf.range(3))
print(out)

输出结果和预期完全一致:

<tf.Tensor: shape=(6, 3), dtype=int32, numpy=
array([[0, 1, 2],
       [0, 2, 1],
       [1, 0, 2],
       [1, 2, 0],
       [2, 0, 1],
       [2, 1, 0]], dtype=int32)>

注意事项

  • 排列总数随输入长度阶乘增长,输入长度过大时会占用大量显存,建议控制输入长度在10以内
  • 输入支持任意dtype的一维张量,输出dtype和输入保持一致
  • 生成的排列顺序和Python标准库itertools.permutations返回顺序完全匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 10:30:04