【问题标题】:Join strings along axis沿轴连接字符串
【发布时间】:2018-10-06 10:09:58
【问题描述】:

假设我有一个 numpy 字符串数组,如下所示:

import numpy as np
print('numpy version:', np.__version__)

a = np.arange(25).reshape(5, 5)
stra = a.astype(np.dtype(str))

print(stra)

输出:

numpy version: 1.15.2
[['0' '1' '2' '3' '4']
 ['5' '6' '7' '8' '9']
 ['10' '11' '12' '13' '14']
 ['15' '16' '17' '18' '19']
 ['20' '21' '22' '23' '24']]

我想沿着给定的轴工作,选择一些元素,然后加入这些字符串。首先我尝试了这个:

print(np.apply_along_axis('|'.join, 1, stra.take([2, 3], 1)))

但结果较长的字符串会被截断以匹配最短的字符串:

['2|3' '7|8' '12|' '17|' '22|']

我当然可以编写自己的循环来获得我想要的输出,但是当单线几乎可以工作时,这样做有点不令人满意。

def join_along_axis(array, indices, axis):        
    if array.ndim == 1:
        return np.array('|'.join(array.take(indices)))

    joined = []        
    # Move axis of interest to end and flatten others to make the loop easy.
    work_arr = np.rollaxis(array, axis, -1)
    shape = work_arr.shape
    new_shape = (np.product(work_arr.shape[:-1]), work_arr.shape[-1])
    work_arr = work_arr.reshape(new_shape)

    for arr in work_arr:
        joined.append('|'.join(arr.take(indices)))

    return np.array(joined).reshape(shape[:-1])

print(join_along_axis(stra, [2, 3], 1))

输出:

['2|3' '7|8' '12|13' '17|18' '22|23']

有没有比我的join_along_axis 函数更巧妙的方法来做到这一点?

为了清楚起见更新:我需要它足够通用,以便在具有任意维数并沿任何选定轴的数组上工作。

【问题讨论】:

  • 我的猜测是apply_along 正在第一行进行测试计算,并使用它来设置结果 dtype。产生较长字符串的行会被截断。查看apply_along 代码。完成所有设置后,它会迭代,为每个“行”调用一次 1d 函数。如果可行,它可能会使任务更简单,但不会更快。

标签: python numpy


【解决方案1】:

我首先尝试使用 apply_along_axis 以您的方式进行操作,但我发现它可能更棘手,apparently NP 没有很好地定义用于处理字符串。

那么列表理解呢?

a =a = np.arange(25).reshape(5, 5)
stra = a.astype(np.dtype(str))
only23 = zip(stra[:,2],stra[:,3])
only23

输出:

[('2', '3'), ('7', '8'), ('12', '13'), ('17', '18'), ('22', ' 23')]

现在让我们进行理解:

[x[0] +'|'+x[1] for x in only23]

输出:

['2|3', '7|8', '12|13', '17|18', '22|23']

你实际上可以让它成为一个单行,我只是不认为它会那么可读

【讨论】:

  • 感谢指向 GitHub 问题的指针。这是一个有用的转换阅读。我应该在我的 OP 中更清楚地知道数组可以有更多维度(或者只有一个,在这种情况下我现有的 join_along_axis 会中断!)如果有更多维度,那么我认为我在函数中所做的重塑将列表理解方法仍然是必要的。我会更新 OP。
【解决方案2】:

从@theshopen 链接的 GitHub 对话中,我似乎可以使用 lambda 来指定我想要的字符串大小。所以这行得通:

lens = np.vectorize(len)
indices = [2, 3]
axis = 1

new_len = lens(stra.take(indices, axis)).sum(1).max() + len(indices) - 1
new_type = '{}{}'.format(stra.dtype.char, new_len)

print(np.apply_along_axis(
    lambda x: np.array('|'.join(x), new_type),
    axis, stra.take(indices, axis)))

【讨论】:

  • 那么速度对比如何?
  • @hpaulj,其实我原来的功能更快。即使我在timeit 循环之外计算new_lennew_type,这种方法仍然需要将近两倍的时间。我认为我的原始函数在可读性方面也很胜一筹,所以也许我会坚持下去!
猜你喜欢
  • 2015-12-15
  • 2021-05-26
  • 1970-01-01
  • 2014-03-31
  • 2018-12-13
  • 1970-01-01
  • 2016-05-25
  • 2020-10-30
  • 2010-10-17
相关资源
最近更新 更多