【发布时间】:2021-08-07 06:23:18
【问题描述】:
各位有经验的朋友,我提出了一种解决算法问题的方法。但是,我发现当数据量增加时,我的方法变得非常耗时。请问有没有更好的方法来解决这个问题?是否可以使用矩阵操作?
问题:
- 假设我们有 1 个
score-matrix和 3 个value-matrix。 - 每个都是
square matrix,大小相同(N*N)。 -
score-matrix中的元素表示两个实体之间的weights。例如,S12表示entity 1和entity 2之间的分数。 (权重仅在大于 0 时才有意义。) -
value-matrix中的元素表示两个实体之间的values。例如,V12表示entity 1和entity 2之间的值。因为我们有 3 个value-matrix,所以我们有 3 个不同的V12。
目标是:我想将values与对应的weights相乘,这样我最终可以输出一个(Nx3)矩阵。
我的解决方案:我解决了这个问题如下。但是,我在这里使用了两个 for 循环,这使得我的程序变得非常耗时。 (例如,当N 很大或3 变为100 时)请问有什么办法可以改进这段代码吗?任何建议或提示将不胜感激。提前谢谢!
# generate sample data
import numpy as np
score_mat = np.random.randint(low=0, high=4, size=(2,2))
value_mat = np.random.randn(3,2,2)
# solve problem
# init the output info
output = np.zeros((2, 3))
# update the output info
for entity_1 in range(2):
# consider meaningful score
entity_others_list = np.where(score_mat[entity_1,:]>0)[0].tolist()
# iterate every other entity
for entity_2 in entity_others_list:
vec = value_mat[:,entity_1,entity_2].copy()
vec *= score_mat[entity_1,entity_2]
output[entity_1] += vec
【问题讨论】:
-
也许您可以将值和分数拆分为矩阵,然后使用矩阵乘法(numpy)?
-
你能跳过 np.where 吗?乘以零比选择非零值更便宜
-
嗨@Stefan。谢谢你们的cmets。我考虑了您的建议,但即使我不使用 np.where,我仍然需要迭代所有列索引。虽然并非所有这些都是非零的。我猜 np.where 限制了范围并提高了效率。如果我做错了什么,请告诉我。
标签: python python-3.x algorithm numpy matrix