如何用Scala/Breeze直接覆写DenseMatrix行中指定列?
如何在Scala/Breeze中实现类似NumPy的矩阵行指定列直接覆写?
我需要把NumPy中直接覆写矩阵某一行指定列的逻辑迁移到Scala/Breeze,目标是用单个表达式完成操作。
NumPy中的实现示例
import numpy as np mat = np.random.normal(size=(2, 5)) print(mat) indexes = np.random.choice(5, replace = False, size = 3) print(indexes) mat[0, [indexes]] = 0 print(mat)
输出:
[[ 0.30389599 0.84549682 -0.38408994 -1.11550844 -0.28496995] [-1.55260273 -0.41368681 -0.40455289 0.13054527 -1.43541557]] [1 3 4] [[ 0.30389599 0. -0.38408994 0. 0. ] [-1.55260273 -0.41368681 -0.40455289 0.13054527 -1.43541557]]
尝试Scala/Breeze时的报错
我尝试直接用类似NumPy的写法,出现类型不匹配错误:
import breeze.linalg.* import breeze.stats.* def main(args: Array[String]): Unit = val mat = DenseMatrix( (-0.25010575, 0.44800905, 0.13285604, 0.34085698, 0.38346101), (-1.97209990, 1.37114368, 1.56601999, -0.13052228, 0.86001178) ) println(mat) val indexes = IndexedSeq(1, 3, 4) println(indexes) mat(0, indexes) = 0.0 println(mat)
错误信息:
-- [E007] Type Mismatch Error: C:\Users\philwalk\workspace\tprf_py\.\rowsliceOverwrite.sc:14:9 14 | mat(0, indexes) = 0 | ^^^^^^^ | Found: (indexes : IndexedSeq[Int]) | Required: Int | | longer explanation available when compiling with `-explain` 1 error found Errors encountered during compilation
当前可行的分步实现
目前我通过三步实现了需求,但不够简洁:
import breeze.linalg.* import breeze.stats.* def main(args: Array[String]): Unit = val mat = DenseMatrix( (-0.25010575, 0.44800905, 0.13285604, 0.34085698, 0.38346101), (-1.97209990, 1.37114368, 1.56601999, -0.13052228, 0.86001178) ) println(mat) val indexes = IndexedSeq(1, 3, 4) println(indexes) var row0 = mat(0, ::).t row0(indexes) := 0.0 mat(0, ::) := row0.t println(mat)
输出:
-0.25010575 0.44800905 0.13285604 0.34085698 0.38346101 -1.9720999 1.37114368 1.56601999 -0.13052228 0.86001178 Vector(1, 3, 4) -0.25010575 0.0 0.13285604 0.0 0.0 -1.9720999 1.37114368 1.56601999 -0.13052228 0.86001178
简洁的单表达式实现方案
可以利用Breeze中Transpose类型的索引支持,直接通过链式索引完成赋值,不需要转置操作:
import breeze.linalg.* import breeze.stats.* def main(args: Array[String]): Unit = val mat = DenseMatrix( (-0.25010575, 0.44800905, 0.13285604, 0.34085698, 0.38346101), (-1.97209990, 1.37114368, 1.56601999, -0.13052228, 0.86001178) ) println(mat) val indexes = IndexedSeq(1, 3, 4) println(indexes) // 单表达式完成指定列置零 mat(0, ::)(indexes) := 0.0 println(mat)
这段代码的输出和分步实现完全一致,且写法更接近NumPy的风格,简洁易读。
内容的提问来源于stack exchange,提问作者philwalk
相关产品推荐
相关产品推荐

