【问题标题】:Create distance matrix using a custom similarity function使用自定义相似度函数创建距离矩阵
【发布时间】:2021-09-28 18:07:12
【问题描述】:

我有一个如下所示的数据框:

data = pd.DataFrame({'id':[1,1,1,2,2,2,3,3,3],
        'age':[20, 21,18,54,23,11, 19, 18,12],
       'experience':[5,4,3,8,2,11,2,8,6]},columns=['id','age','experience'])

   id  age experience
0   1   20  5
1   1   21  4
2   1   18  3
3   2   54  8
4   2   23  2
5   2   11  11
6   3   19  2
7   3   18  8
8   3   12  6

我正在使用一个名为 dtw_path 的自定义距离函数来计算元组之间的距离。我不会详细讨论该函数如何将距离计算为一个复杂的过程,但它只是输出元组之间的标量距离值。

元组的形成方式如下:

data['age_exp'] = data[['age', 'experience']].apply(tuple, axis=1)

    id  age experience  age_exp
0   1   20   5          (20, 5)
1   1   21   4          (21, 4)
2   1   18   3          (18, 3)
3   2   54   8          (54, 8)
4   2   23   2          (23, 2)
5   2   11   11         (11, 11)
6   3   19   2          (19, 2)
7   3   18   8          (18, 8)
8   3   12   6          (12, 6)

所以对于上述数据框,如果我需要计算 id 1 和 id 2 之间的距离,我会按如下方式计算距离:

data1 = data[data['id']==1]
data1 = np.array(data1['age_exp'].tolist())
data1

array([[20,  5],
       [21,  4],
       [18,  3]])

data2 = data[data['id']==2]
data2 = np.array(data2['age_exp'].tolist())
data2

array([[54,  8],
       [23,  2],
       [11, 11]])

dtw_path(data1,data2)[1]

1.5

我需要帮助的是如何遍历数据框并为 id 列创建距离矩阵,即类似这样的东西

     1    2     3
1    0    1.5   2          
2    1.5  0     2.3
3    2    2.3   0

【问题讨论】:

  • dtw_path 到底是什么?

标签: python pandas numpy matrix distance


【解决方案1】:

您的问题不清楚dtw_path 是什么。我在这里使用了tslearn.metrics.dtw_path,这给了我不同的结果。然而,基本原理应该是一样的。

让我们先重塑一下原始数据框:

data2 = (data.groupby('id')
             .apply(lambda x: np.array(list(zip(x['age'], x['experience']))))
        ).to_frame()
                               0
id                              
1    [[20, 5], [21, 4], [18, 3]]
2   [[54, 8], [23, 2], [11, 11]]
3    [[19, 2], [18, 8], [12, 6]]

注意。下一步需要二维(DataFrame),因此.to_frame()

然后,使用scipy.spatial.distance.pdist 应用您的dtw_path 函数,该函数可以使用参数metric 获取任意距离函数,并仅保留输出的第二个元素。最后,使用scipy.spatial.distance.squareform 将输出重塑为方阵:

squareform(pdist(data2, metric=lambda x,y: dtw_path(x[0], y[0])[1]))

输出:

array([[ 0.        , 35.86084215,  8.94427191],
       [35.86084215,  0.        , 36.7151195 ],
       [ 8.94427191, 36.7151195 ,  0.        ]])

【讨论】:

  • 这是完美的,是的,您的 dtw_path 的值是正确的,我只是为了说明目的而更改了它。
猜你喜欢
  • 2019-03-16
  • 1970-01-01
  • 2014-03-20
  • 1970-01-01
  • 1970-01-01
  • 2023-03-07
  • 1970-01-01
  • 2015-06-11
  • 1970-01-01
相关资源
最近更新 更多