【问题标题】:query dataframe column on array values在数组值上查询数据框列
【发布时间】:2019-01-03 16:31:36
【问题描述】:
traj0
Out[52]: 
         state         action  reward
0   [1.0, 4.0, 6.0]     3.0     4.0
1   [4.0, 6.0, 11.0]    4.0     5.0
2   [6.0, 7.0, 3.0]     3.0    22.0
3   [3.0, 3.0, 2.0]     1.0    10.0
4   [2.0, 9.0, 5.0]     2.0     2.0

假设我有一个看起来像这样的 pandas 数据框,其中状态列作为其条目,3 元素 numpy 数组。

如何在此处查询状态为np.array([3.0,3.0,2.0]) 的行?

我知道traj0.query("state == '[3.0,3.0,2.0]'") 有效,我知道。但我不想在查询中硬编码数组值。

我正在寻找类似的东西

x = np.array([3.0,3.0,2.0])
traj0.query('state ==' + x)

=============

这不是一个重复的问题,因为我之前的问题pandas query with a column consisting of array entries 仅适用于每个数组中只有一个值的情况。在这里我正在寻找数组是否有多个值。

【问题讨论】:

  • 为什么不将它们存储在单独的列中:state_1state_2state_3
  • 我们如何定义状态是多个值的向量。我必须为这个项目。

标签: python arrays pandas numpy dataframe


【解决方案1】:
import numpy as np
import pandas as pd

df = pd.DataFrame([[np.array([1.0, 4.0, 6.0]), 3.0, 4.0],
              [np.array([4.0, 6.0, 11.0]), 4.0, 5.0],
              [np.array([6.0, 7.0, 3.0]), 3.0, 22.0],
              [np.array([3.0, 3.0, 2.0]), 1.0, 10.0],
              [np.array([2.0, 9.0, 5.0]), 2.0, 2.0]
             ], columns=['state','action','reward'])

x = str(np.array([3.0, 3.0, 2.0]))
df[df.state.astype(str) == x]

// to use pd.query
df['state_str'] = df.state.astype(str)
df.query("state_str == '{}'".format(x))

输出

    state           action  reward
3   [3.0, 3.0, 2.0] 1.0     10.0

【讨论】:

  • 注意字符串转换是昂贵的。一般来说,NumPy 数组不需要它。
  • @jpp 我同意,我真的很喜欢你的解决方案。但是,如果 @Glassjawed 想要使用 pd.query,这是一种方法。
【解决方案2】:

您可以使用 df.loc 和使用 numpy.array_equal 的 lambda 函数来做到这一点:

x = [1., 4., 6.]
traj0.loc[df.state.apply(lambda a: np.array_equal(a, x))]

基本上,这会检查state 列的每个元素是否与x 等效,并仅返回与该列匹配的那些行。

示例

df = pd.DataFrame(data={'state': [[1., 4., 6.], [4., 5., 6.]],
                        'value': [5, 6]})
print(df.loc[df.state.apply(lambda a: np.array_equal(a, x))])

             state  value
0  [1.0, 4.0, 6.0]      5

【讨论】:

  • 请注意,这等效于 for 循环。它没有利用 NumPy 矢量化操作。
【解决方案3】:

最好在这里使用pd.DataFrame.query。您可以执行矢量化比较,然后使用布尔索引:

x = [3, 3, 2]
mask = (np.array(df['state'].values.tolist()) == x).all(1)

res = df[mask]

print(res)

             state  action  reward
3  [3.0, 3.0, 2.0]     1.0    10.0

一般来说,您不应该在 Pandas 系列中存储列表或数组。这是低效的并且消除了直接矢量化操作的可能性。在这里,我们必须显式转换为 NumPy 数组以进行简单比较。

【讨论】:

  • 我不希望状态值被硬编码。那么如果我说 x=[3,3,2]...我该如何使用呢?
猜你喜欢
  • 2021-04-27
  • 2017-12-18
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2017-04-04
  • 1970-01-01
  • 2022-11-25
  • 1970-01-01
相关资源
最近更新 更多